@@ -63,6 +63,17 @@ def forward(self, logits, temperature, seed, top_p):
6363 return self .head (logits , temperature = temperature , seed = seed , top_p = top_p )
6464
6565
66+ class TopKSampleModel (nn .Module ):
67+ """SamplingHead with temperature, seed, and top_k as runtime inputs."""
68+
69+ def __init__ (self ):
70+ super ().__init__ ()
71+ self .head = SamplingHead (_LogitsPassthrough ())
72+
73+ def forward (self , logits , temperature , seed , top_k ):
74+ return self .head (logits , temperature = temperature , seed = seed , top_k = top_k )
75+
76+
6677def _ref_gumbel_max (logits : torch .Tensor , temperature : float , seed : int ):
6778 """Independent Gumbel-max reference using the same torch RNG as the op."""
6879 gen = torch .Generator ().manual_seed (seed )
@@ -76,11 +87,21 @@ def _tv_distance(p: torch.Tensor, q: torch.Tensor) -> float:
7687 return 0.5 * torch .abs (p - q ).sum ().item ()
7788
7889
79- def _sample (logits , temperature , seed : Optional [int ], top_p : float = 1.0 ):
90+ def _sample (
91+ logits ,
92+ temperature ,
93+ seed : Optional [int ],
94+ top_p : float = 1.0 ,
95+ top_k : Optional [int ] = None ,
96+ ):
8097 t = torch .tensor (float (temperature ))
8198 s = None if seed is None else torch .tensor (int (seed ), dtype = torch .int64 )
8299 p = torch .tensor (float (top_p )) # 1.0 = off
83- return torch .ops .mlx .sample (logits , t , p , s )
100+ k = torch .tensor (
101+ torch .iinfo (torch .int64 ).max if top_k is None else int (top_k ),
102+ dtype = torch .int64 ,
103+ )
104+ return torch .ops .mlx .sample (logits , t , k , p , s )
84105
85106
86107class TestSampleOp (unittest .TestCase ):
@@ -142,6 +163,33 @@ def test_top_p_one_keeps_all(self):
142163 tokens = _sample (base .expand (20000 , 4 ), 1.0 , seed = 0 , top_p = 1.0 )
143164 self .assertTrue ((tokens == 3 ).any ())
144165
166+ def test_top_k_restricts_to_top_k (self ):
167+ # Non-sorted probs [0.15, 0.5, 0.05, 0.3]; top_k=2 keeps {1,3}.
168+ base = torch .log (torch .tensor ([0.15 , 0.5 , 0.05 , 0.3 ]))
169+ tokens = _sample (base .expand (5000 , 4 ), 1.0 , seed = 0 , top_k = 2 )
170+ self .assertTrue (torch .isin (tokens , torch .tensor ([1 , 3 ])).all ())
171+ self .assertEqual (set (tokens .tolist ()), {1 , 3 })
172+
173+ def test_top_k_default_keeps_all (self ):
174+ # top_k=None -> no filtering; the tail token (index 3) is reachable.
175+ base = torch .log (torch .tensor ([0.5 , 0.3 , 0.15 , 0.05 ]))
176+ tokens = _sample (base .expand (20000 , 4 ), 1.0 , seed = 0 , top_k = None )
177+ self .assertTrue ((tokens == 3 ).any ())
178+
179+ def test_top_k_clips_to_vocab_size (self ):
180+ # top_k > vocab is clipped to vocab size, so every token is reachable.
181+ base = torch .log (torch .tensor ([0.5 , 0.3 , 0.15 , 0.05 ]))
182+ tokens = _sample (base .expand (20000 , 4 ), 1.0 , seed = 0 , top_k = 999 )
183+ self .assertEqual (set (tokens .tolist ()), {0 , 1 , 2 , 3 })
184+
185+ def test_top_k_and_top_p_compose (self ):
186+ # top_k is applied before top_p, so top_p sees renormalized top-k probs.
187+ # top_k=3 -> [0.526, 0.316, 0.158]; top_p=0.83 keeps {0,1}.
188+ base = torch .log (torch .tensor ([0.5 , 0.3 , 0.15 , 0.05 ]))
189+ tokens = _sample (base .expand (5000 , 4 ), 1.0 , seed = 0 , top_p = 0.83 , top_k = 3 )
190+ self .assertTrue (torch .isin (tokens , torch .tensor ([0 , 1 ])).all ())
191+ self .assertEqual (set (tokens .tolist ()), {0 , 1 })
192+
145193
146194class TestSampleExport (unittest .TestCase ):
147195 """Runtime-input semantics that survive export: temperature and seed stay
@@ -218,6 +266,24 @@ def test_top_p_end_to_end(self):
218266 (token ,) = load_tensors_from_bin (out_bin )
219267 self .assertIn (int (token ), {0 , 1 , 2 }) # tail token (index 3) excluded
220268
269+ def test_top_k_end_to_end (self ):
270+ # On-device top-k: probs [0.5,0.3,0.15,0.05], top_k=2 -> token in {0,1}.
271+ logits = torch .log (torch .tensor ([0.5 , 0.3 , 0.15 , 0.05 ])).view (1 , 1 , 4 )
272+ inputs = (
273+ logits ,
274+ torch .tensor (1.0 ),
275+ torch .tensor (0 , dtype = torch .int64 ),
276+ torch .tensor (2 , dtype = torch .int64 ),
277+ )
278+ tmp = Path (self ._tmp )
279+ pte , in_bin , out_bin = tmp / "topk.pte" , tmp / "in.bin" , tmp / "out.bin"
280+ export_model_to_pte (TopKSampleModel (), inputs , pte )
281+ save_tensors_to_bin (list (inputs ), in_bin )
282+
283+ self .assertTrue (run_cpp_test_runner (pte , in_bin , out_bin ))
284+ (token ,) = load_tensors_from_bin (out_bin )
285+ self .assertIn (int (token ), {0 , 1 }) # tail tokens excluded
286+
221287
222288if __name__ == "__main__" :
223289 unittest .main ()
0 commit comments