From c944a5a73bbef6af87f5af38947a0502ea52faf5 Mon Sep 17 00:00:00 2001 From: goutamadwant Date: Sat, 27 Jun 2026 10:34:20 -0700 Subject: [PATCH 1/3] Add top-k support to MLX sample --- backends/mlx/custom_ops.py | 12 +++++-- backends/mlx/llm/sampling.py | 9 ++--- backends/mlx/ops.py | 44 +++++++++++++++++++++++- backends/mlx/test/test_ops.py | 44 ++++++++++++++++++++++++ backends/mlx/test/test_sample.py | 59 ++++++++++++++++++++++++++++++-- 5 files changed, 159 insertions(+), 9 deletions(-) diff --git a/backends/mlx/custom_ops.py b/backends/mlx/custom_ops.py index 17c07097f70..7d53c4ebe5b 100644 --- a/backends/mlx/custom_ops.py +++ b/backends/mlx/custom_ops.py @@ -399,9 +399,11 @@ def sample( temperature: Tensor, top_p: Tensor, seed: Optional[Tensor] = None, + top_k: Optional[Tensor] = None, ) -> Tensor: """ - Gumbel-max sampling from softmax(logits / temperature), with top-p (nucleus). + Gumbel-max sampling from softmax(logits / temperature), with top-p (nucleus) + and optional top-k filtering. logits: [B, vocab] temperature: scalar float tensor (runtime input). temperature <= 0 is greedy: return argmax(logits) directly (matches the device, @@ -411,6 +413,8 @@ def sample( seed: scalar int tensor or None - tensor -> deterministic, keyed RNG (random::key(seed)) - None -> MLX global KeySequence (non-deterministic) + top_k: optional scalar int tensor. None keeps every token; otherwise + sampling is restricted to the k most likely tokens. -> token_id: [B] int64 Host/CPU reference used for export (shape/meta) and distributional checks @@ -424,6 +428,10 @@ def sample( scaled = logits.float() / temperature probs = torch.softmax(scaled, dim=-1) s_probs, _ = torch.sort(probs, dim=-1, descending=True) + if top_k is not None: + k = int(top_k.item()) + kth = s_probs[..., k - 1 : k] + scaled = torch.where(probs >= kth, scaled, scaled.new_tensor(float("-inf"))) cum = torch.cumsum(s_probs, dim=-1) keep = (cum - s_probs) <= top_p thresh = torch.where(keep, s_probs, s_probs.new_tensor(float("inf"))).amin( @@ -440,5 +448,5 @@ def sample( @torch.library.register_fake("mlx::sample") -def sample_fake(logits, temperature, top_p, seed=None): +def sample_fake(logits, temperature, top_p, seed=None, top_k=None): return logits.new_empty(logits.shape[:-1], dtype=torch.long) diff --git a/backends/mlx/llm/sampling.py b/backends/mlx/llm/sampling.py index a059e5cff08..eb32635e2c6 100644 --- a/backends/mlx/llm/sampling.py +++ b/backends/mlx/llm/sampling.py @@ -19,7 +19,8 @@ class SamplingHead(nn.Module): temperature: scalar float tensor, e.g. torch.tensor(0.8). Must be >= 0; temperature=0 is greedy (returns argmax, no division). - top_k: not implemented yet (reserved); must be None. + top_k: scalar int tensor or int; keeps only the k most likely tokens. + None disables top-k filtering. top_p: scalar float tensor in (0, 1] for nucleus sampling. top_p=1.0 (the default) keeps every token, i.e. no filtering. Pass it as a runtime input to tune per request. @@ -31,10 +32,10 @@ def __init__(self, model: nn.Module): self.model = model def forward(self, *args, temperature, top_k=None, top_p=1.0, seed=None, **kwargs): - if top_k is not None: - raise NotImplementedError("top_k sampling is not implemented") logits = self.model(*args, **kwargs) # [B, S, vocab] last = logits[:, -1, :] # [B, vocab] if not isinstance(top_p, torch.Tensor): top_p = torch.tensor(float(top_p)) - return torch.ops.mlx.sample(last, temperature, top_p, seed) + if top_k is not None and not isinstance(top_k, torch.Tensor): + top_k = torch.tensor(int(top_k), dtype=torch.int64) + return torch.ops.mlx.sample(last, temperature, top_p, seed, top_k) diff --git a/backends/mlx/ops.py b/backends/mlx/ops.py index 86e322a16e7..6a35ca7cc87 100644 --- a/backends/mlx/ops.py +++ b/backends/mlx/ops.py @@ -24,6 +24,7 @@ emit_lifted_constant, emit_quantized_biases, emit_shape, + emit_sub_int, parse_dequant_node, to_mlx_qparams, torch_dtype_to_scalar_type, @@ -3528,10 +3529,11 @@ def _sample_handler(P: MLXProgramBuilder, n: Node) -> Slot: skipping the sampling chain (so 0 is exact, not the small-epsilon approx). """ args = P.args(n) - require_args(args, 3, 4, "mlx.sample") + require_args(args, 3, 5, "mlx.sample") require_kwargs(P.kwargs(n), set(), "mlx.sample") logits, temperature, top_p = args[0], args[1], args[2] seed = args[3] if len(args) > 3 and args[3] is not None else None + top_k = args[4] if len(args) > 4 and args[4] is not None else None temp_dt = n.args[1].meta["val"].dtype out = P.make_or_get_slot(n) @@ -3674,6 +3676,46 @@ def emit_sample(): out=P.slot_to_tid(drop), ) ) + if top_k is not None: + _, top_k_val = P.make_tmp_value_slot() + P.emit(ItemIntNode(x=P.slot_to_tid(top_k), out=P.slot_to_vid(top_k_val))) + top_k_iov = P.to_int_or_vid(top_k_val) + top_k_index = emit_sub_int(P, top_k_iov, IntOrVid.from_literal(1)) + _, top_k_thresh = P.make_tmp_slot() + P.emit( + TakeNode( + x=P.slot_to_tid(sorted_p), + index=( + IntOrVidOrTid.from_vid(top_k_index.vid) + if top_k_index.is_vid + else IntOrVidOrTid.from_literal(top_k_index.literal) + ), + out=P.slot_to_tid(top_k_thresh), + axis=-1, + ) + ) + P.emit( + ExpandDimsNode( + x=P.slot_to_tid(top_k_thresh), + out=P.slot_to_tid(top_k_thresh), + axis=-1, + ) + ) + _, drop_k = P.make_tmp_slot() + P.emit( + LessNode( + a=P.slot_to_tid(probs), + b=P.slot_to_tid(top_k_thresh), + out=P.slot_to_tid(drop_k), + ) + ) + P.emit( + LogicalOrNode( + a=P.slot_to_tid(drop), + b=P.slot_to_tid(drop_k), + out=P.slot_to_tid(drop), + ) + ) neg_inf = emit_lifted_constant(P, float("-inf"), torch.float32) # masked = where(drop, -inf, scaled); then add gumbel noise in place. _, masked = P.make_tmp_slot() diff --git a/backends/mlx/test/test_ops.py b/backends/mlx/test/test_ops.py index afd4f276dde..bde3f495343 100644 --- a/backends/mlx/test/test_ops.py +++ b/backends/mlx/test/test_ops.py @@ -7676,6 +7676,17 @@ def forward(self, logits, temperature, seed, top_p): return self.head(logits, temperature=temperature, seed=seed, top_p=top_p) +class TopKSampleModel(nn.Module): + """SamplingHead with temperature, seed, and top_k as runtime inputs.""" + + def __init__(self): + super().__init__() + self.head = SamplingHead(_LogitsPassthrough()) + + def forward(self, logits, temperature, seed, top_k): + return self.head(logits, temperature=temperature, seed=seed, top_k=top_k) + + @register_test class SampleSeededTest(OpTestCase): """Seeded sample lowers to one MLX segment; seed threads in via ItemIntNode.""" @@ -7756,6 +7767,39 @@ def create_inputs(self) -> Tuple[torch.Tensor, ...]: ) +@register_test +class SampleTopKTest(OpTestCase): + """Top-k sample emits the extra threshold and combined mask nodes.""" + + name = "sample_top_k" + skip_comparison = True # sampling RNG is not host/device bit-identical + expected_node_counts = { + "IfNode": 1, + "RandomBitsNode": 1, + "ArgmaxNode": 2, + "ItemIntNode": 3, # seed + top_k + temperature>0 condition + "SoftmaxNode": 1, + "SortNode": 1, + "CumsumNode": 1, + "MinNode": 1, + "TakeNode": 1, + "ExpandDimsNode": 1, + "LogicalOrNode": 1, + "WhereNode": 2, + } + + def create_model(self) -> nn.Module: + return TopKSampleModel() + + def create_inputs(self) -> Tuple[torch.Tensor, ...]: + return ( + torch.randn(1, 4, 256), + torch.tensor(0.8), + torch.tensor(0, dtype=torch.int64), + torch.tensor(2, dtype=torch.int64), + ) + + @register_test class SampleGreedyTest(OpTestCase): """Greedy argmax(logits) is bit-exact host/device, so verify the token with the diff --git a/backends/mlx/test/test_sample.py b/backends/mlx/test/test_sample.py index ddb0734d13b..276677d9993 100644 --- a/backends/mlx/test/test_sample.py +++ b/backends/mlx/test/test_sample.py @@ -63,6 +63,17 @@ def forward(self, logits, temperature, seed, top_p): return self.head(logits, temperature=temperature, seed=seed, top_p=top_p) +class TopKSampleModel(nn.Module): + """SamplingHead with temperature, seed, and top_k as runtime inputs.""" + + def __init__(self): + super().__init__() + self.head = SamplingHead(_LogitsPassthrough()) + + def forward(self, logits, temperature, seed, top_k): + return self.head(logits, temperature=temperature, seed=seed, top_k=top_k) + + def _ref_gumbel_max(logits: torch.Tensor, temperature: float, seed: int): """Independent Gumbel-max reference using the same torch RNG as the op.""" gen = torch.Generator().manual_seed(seed) @@ -76,11 +87,18 @@ def _tv_distance(p: torch.Tensor, q: torch.Tensor) -> float: return 0.5 * torch.abs(p - q).sum().item() -def _sample(logits, temperature, seed: Optional[int], top_p: float = 1.0): +def _sample( + logits, + temperature, + seed: Optional[int], + top_p: float = 1.0, + top_k: Optional[int] = None, +): t = torch.tensor(float(temperature)) s = None if seed is None else torch.tensor(int(seed), dtype=torch.int64) p = torch.tensor(float(top_p)) # 1.0 = off - return torch.ops.mlx.sample(logits, t, p, s) + k = None if top_k is None else torch.tensor(int(top_k), dtype=torch.int64) + return torch.ops.mlx.sample(logits, t, p, s, k) class TestSampleOp(unittest.TestCase): @@ -142,6 +160,25 @@ def test_top_p_one_keeps_all(self): tokens = _sample(base.expand(20000, 4), 1.0, seed=0, top_p=1.0) self.assertTrue((tokens == 3).any()) + def test_top_k_restricts_to_top_k(self): + # probs [0.5, 0.3, 0.15, 0.05]; top_k=2 keeps {0,1}. + base = torch.log(torch.tensor([0.5, 0.3, 0.15, 0.05])) + tokens = _sample(base.expand(5000, 4), 1.0, seed=0, top_k=2) + self.assertTrue(torch.isin(tokens, torch.tensor([0, 1])).all()) + self.assertEqual(set(tokens.tolist()), {0, 1}) + + def test_top_k_none_keeps_all(self): + # top_k=None -> no filtering; the tail token (index 3) is reachable. + base = torch.log(torch.tensor([0.5, 0.3, 0.15, 0.05])) + tokens = _sample(base.expand(20000, 4), 1.0, seed=0, top_k=None) + self.assertTrue((tokens == 3).any()) + + def test_top_k_and_top_p_compose(self): + # top_p=0.7 keeps {0,1}; top_k=1 intersects that to {0}. + base = torch.log(torch.tensor([0.5, 0.3, 0.15, 0.05])) + tokens = _sample(base.expand(5000, 4), 1.0, seed=0, top_p=0.7, top_k=1) + self.assertEqual(set(tokens.tolist()), {0}) + class TestSampleExport(unittest.TestCase): """Runtime-input semantics that survive export: temperature and seed stay @@ -218,6 +255,24 @@ def test_top_p_end_to_end(self): (token,) = load_tensors_from_bin(out_bin) self.assertIn(int(token), {0, 1, 2}) # tail token (index 3) excluded + def test_top_k_end_to_end(self): + # On-device top-k: probs [0.5,0.3,0.15,0.05], top_k=2 -> token in {0,1}. + logits = torch.log(torch.tensor([0.5, 0.3, 0.15, 0.05])).view(1, 1, 4) + inputs = ( + logits, + torch.tensor(1.0), + torch.tensor(0, dtype=torch.int64), + torch.tensor(2, dtype=torch.int64), + ) + tmp = Path(self._tmp) + pte, in_bin, out_bin = tmp / "topk.pte", tmp / "in.bin", tmp / "out.bin" + export_model_to_pte(TopKSampleModel(), inputs, pte) + save_tensors_to_bin(list(inputs), in_bin) + + self.assertTrue(run_cpp_test_runner(pte, in_bin, out_bin)) + (token,) = load_tensors_from_bin(out_bin) + self.assertIn(int(token), {0, 1}) # tail tokens excluded + if __name__ == "__main__": unittest.main() From 7ede2d50e8954c07ffc37185b71784837d19c918 Mon Sep 17 00:00:00 2001 From: goutamadwant Date: Sat, 27 Jun 2026 22:33:08 -0700 Subject: [PATCH 2/3] Address MLX top-k sampling review --- backends/mlx/custom_ops.py | 24 +++--- backends/mlx/llm/sampling.py | 9 ++- backends/mlx/ops.py | 127 +++++++++++++++++++------------ backends/mlx/test/test_ops.py | 26 ++++--- backends/mlx/test/test_sample.py | 31 +++++--- 5 files changed, 136 insertions(+), 81 deletions(-) diff --git a/backends/mlx/custom_ops.py b/backends/mlx/custom_ops.py index 7d53c4ebe5b..5605b59c543 100644 --- a/backends/mlx/custom_ops.py +++ b/backends/mlx/custom_ops.py @@ -397,24 +397,24 @@ def gather_qmm_fake( def sample( logits: Tensor, temperature: Tensor, + top_k: Tensor, top_p: Tensor, seed: Optional[Tensor] = None, - top_k: Optional[Tensor] = None, ) -> Tensor: """ - Gumbel-max sampling from softmax(logits / temperature), with top-p (nucleus) - and optional top-k filtering. + Gumbel-max sampling from softmax(logits / temperature), with top-k and + top-p (nucleus) filtering. logits: [B, vocab] temperature: scalar float tensor (runtime input). temperature <= 0 is greedy: return argmax(logits) directly (matches the device, which branches on temperature > 0). + top_k: scalar int tensor. It is clipped to the vocab size; using the + max int default keeps every token. top_p: scalar float tensor in (0, 1]. top_p=1.0 keeps every token, i.e. it is off. seed: scalar int tensor or None - tensor -> deterministic, keyed RNG (random::key(seed)) - None -> MLX global KeySequence (non-deterministic) - top_k: optional scalar int tensor. None keeps every token; otherwise - sampling is restricted to the k most likely tokens. -> token_id: [B] int64 Host/CPU reference used for export (shape/meta) and distributional checks @@ -426,12 +426,16 @@ def sample( return torch.argmax(logits, dim=-1) # whole chain in fp32 to match the lowered graph (bf16 sums mis-rank ties). scaled = logits.float() / temperature + + k = min(int(top_k.item()), scaled.shape[-1]) + s_scaled, _ = torch.sort(scaled, dim=-1, descending=True) + kth = s_scaled[..., k - 1 : k] + scaled = torch.where(scaled >= kth, scaled, scaled.new_tensor(float("-inf"))) + + # Apply top-p after top-k so the probabilities are renormalized over the + # top-k subset. probs = torch.softmax(scaled, dim=-1) s_probs, _ = torch.sort(probs, dim=-1, descending=True) - if top_k is not None: - k = int(top_k.item()) - kth = s_probs[..., k - 1 : k] - scaled = torch.where(probs >= kth, scaled, scaled.new_tensor(float("-inf"))) cum = torch.cumsum(s_probs, dim=-1) keep = (cum - s_probs) <= top_p thresh = torch.where(keep, s_probs, s_probs.new_tensor(float("inf"))).amin( @@ -448,5 +452,5 @@ def sample( @torch.library.register_fake("mlx::sample") -def sample_fake(logits, temperature, top_p, seed=None, top_k=None): +def sample_fake(logits, temperature, top_k, top_p, seed=None): return logits.new_empty(logits.shape[:-1], dtype=torch.long) diff --git a/backends/mlx/llm/sampling.py b/backends/mlx/llm/sampling.py index eb32635e2c6..fb94ceb7d27 100644 --- a/backends/mlx/llm/sampling.py +++ b/backends/mlx/llm/sampling.py @@ -20,7 +20,8 @@ class SamplingHead(nn.Module): temperature: scalar float tensor, e.g. torch.tensor(0.8). Must be >= 0; temperature=0 is greedy (returns argmax, no division). top_k: scalar int tensor or int; keeps only the k most likely tokens. - None disables top-k filtering. + None uses the max int default, which is clipped to the vocab + size and keeps every token. top_p: scalar float tensor in (0, 1] for nucleus sampling. top_p=1.0 (the default) keeps every token, i.e. no filtering. Pass it as a runtime input to tune per request. @@ -36,6 +37,8 @@ def forward(self, *args, temperature, top_k=None, top_p=1.0, seed=None, **kwargs last = logits[:, -1, :] # [B, vocab] if not isinstance(top_p, torch.Tensor): top_p = torch.tensor(float(top_p)) - if top_k is not None and not isinstance(top_k, torch.Tensor): + if top_k is None: + top_k = torch.tensor(torch.iinfo(torch.int64).max, dtype=torch.int64) + elif not isinstance(top_k, torch.Tensor): top_k = torch.tensor(int(top_k), dtype=torch.int64) - return torch.ops.mlx.sample(last, temperature, top_p, seed, top_k) + return torch.ops.mlx.sample(last, temperature, top_k, top_p, seed) diff --git a/backends/mlx/ops.py b/backends/mlx/ops.py index 6a35ca7cc87..986b713a5ff 100644 --- a/backends/mlx/ops.py +++ b/backends/mlx/ops.py @@ -24,7 +24,6 @@ emit_lifted_constant, emit_quantized_biases, emit_shape, - emit_sub_int, parse_dequant_node, to_mlx_qparams, torch_dtype_to_scalar_type, @@ -3529,11 +3528,10 @@ def _sample_handler(P: MLXProgramBuilder, n: Node) -> Slot: skipping the sampling chain (so 0 is exact, not the small-epsilon approx). """ args = P.args(n) - require_args(args, 3, 5, "mlx.sample") + require_args(args, 4, 5, "mlx.sample") require_kwargs(P.kwargs(n), set(), "mlx.sample") - logits, temperature, top_p = args[0], args[1], args[2] - seed = args[3] if len(args) > 3 and args[3] is not None else None - top_k = args[4] if len(args) > 4 and args[4] is not None else None + logits, temperature, top_k, top_p = args[0], args[1], args[2], args[3] + seed = args[4] if len(args) > 4 and args[4] is not None else None temp_dt = n.args[1].meta["val"].dtype out = P.make_or_get_slot(n) @@ -3614,8 +3612,84 @@ def emit_sample(): ) ) scaled = logits_f + neg_inf = emit_lifted_constant(P, float("-inf"), torch.float32) + + # Top-k first, on scaled logits. Clip k to vocab size so the default + # max-int sentinel selects every token. + vocab_size = int(n.args[0].meta["val"].shape[-1]) + vocab = emit_lifted_constant(P, vocab_size, torch.int64) + _, clipped_top_k = P.make_tmp_slot() + P.emit( + MinimumNode( + a=P.slot_to_tid(top_k), + b=P.slot_to_tid(vocab), + out=P.slot_to_tid(clipped_top_k), + ) + ) + _, top_k_val = P.make_tmp_value_slot() + P.emit( + ItemIntNode( + x=P.slot_to_tid(clipped_top_k), out=P.slot_to_vid(top_k_val) + ) + ) + _, top_k_index = P.make_tmp_value_slot() + P.emit( + SubtractIntNode( + a=P.to_int_or_vid(top_k_val), + b=IntOrVid.from_literal(1), + out=P.slot_to_vid(top_k_index), + ) + ) - # top-p nucleus mask; SortNode is ascending-only, so sort -probs for descending. + _, sorted_scaled = P.make_tmp_slot() + P.emit(NegNode(x=P.slot_to_tid(scaled), out=P.slot_to_tid(sorted_scaled))) + P.emit( + SortNode( + x=P.slot_to_tid(sorted_scaled), + out=P.slot_to_tid(sorted_scaled), + axis=-1, + ) + ) + P.emit( + NegNode(x=P.slot_to_tid(sorted_scaled), out=P.slot_to_tid(sorted_scaled)) + ) + _, top_k_thresh = P.make_tmp_slot() + P.emit( + TakeNode( + x=P.slot_to_tid(sorted_scaled), + index=P.to_int_or_vid_or_tid(top_k_index), + out=P.slot_to_tid(top_k_thresh), + axis=-1, + ) + ) + P.emit( + ExpandDimsNode( + x=P.slot_to_tid(top_k_thresh), + out=P.slot_to_tid(top_k_thresh), + axis=-1, + ) + ) + _, drop_k = P.make_tmp_slot() + P.emit( + LessNode( + a=P.slot_to_tid(scaled), + b=P.slot_to_tid(top_k_thresh), + out=P.slot_to_tid(drop_k), + ) + ) + _, top_k_scaled = P.make_tmp_slot() + P.emit( + WhereNode( + condition=P.slot_to_tid(drop_k), + x=P.slot_to_tid(neg_inf), + y=P.slot_to_tid(scaled), + out=P.slot_to_tid(top_k_scaled), + ) + ) + scaled = top_k_scaled + + # Top-p nucleus mask on probabilities renormalized over the top-k set. + # SortNode is ascending-only, so sort -probs for descending. # probs is read twice (neg_p below and the drop comparison), keep separate. _, probs = P.make_tmp_slot() P.emit(SoftmaxNode(x=P.slot_to_tid(scaled), out=P.slot_to_tid(probs), axis=-1)) @@ -3676,47 +3750,6 @@ def emit_sample(): out=P.slot_to_tid(drop), ) ) - if top_k is not None: - _, top_k_val = P.make_tmp_value_slot() - P.emit(ItemIntNode(x=P.slot_to_tid(top_k), out=P.slot_to_vid(top_k_val))) - top_k_iov = P.to_int_or_vid(top_k_val) - top_k_index = emit_sub_int(P, top_k_iov, IntOrVid.from_literal(1)) - _, top_k_thresh = P.make_tmp_slot() - P.emit( - TakeNode( - x=P.slot_to_tid(sorted_p), - index=( - IntOrVidOrTid.from_vid(top_k_index.vid) - if top_k_index.is_vid - else IntOrVidOrTid.from_literal(top_k_index.literal) - ), - out=P.slot_to_tid(top_k_thresh), - axis=-1, - ) - ) - P.emit( - ExpandDimsNode( - x=P.slot_to_tid(top_k_thresh), - out=P.slot_to_tid(top_k_thresh), - axis=-1, - ) - ) - _, drop_k = P.make_tmp_slot() - P.emit( - LessNode( - a=P.slot_to_tid(probs), - b=P.slot_to_tid(top_k_thresh), - out=P.slot_to_tid(drop_k), - ) - ) - P.emit( - LogicalOrNode( - a=P.slot_to_tid(drop), - b=P.slot_to_tid(drop_k), - out=P.slot_to_tid(drop), - ) - ) - neg_inf = emit_lifted_constant(P, float("-inf"), torch.float32) # masked = where(drop, -inf, scaled); then add gumbel noise in place. _, masked = P.make_tmp_slot() P.emit( diff --git a/backends/mlx/test/test_ops.py b/backends/mlx/test/test_ops.py index bde3f495343..bd3a167f56a 100644 --- a/backends/mlx/test/test_ops.py +++ b/backends/mlx/test/test_ops.py @@ -7697,12 +7697,14 @@ class SampleSeededTest(OpTestCase): "IfNode": 1, # temperature==0 greedy branch "RandomBitsNode": 1, "ArgmaxNode": 2, # sampling branch + greedy branch - "ItemIntNode": 2, # seed + temperature>0 condition + "ItemIntNode": 3, # seed + top_k + temperature>0 condition "SoftmaxNode": 1, # top-p nucleus chain - "SortNode": 1, + "SortNode": 2, # top-k threshold + top-p nucleus chain "CumsumNode": 1, "MinNode": 1, - "WhereNode": 2, + "TakeNode": 1, + "ExpandDimsNode": 1, + "WhereNode": 3, } def create_model(self) -> nn.Module: @@ -7726,7 +7728,7 @@ class SampleUnseededTest(OpTestCase): "IfNode": 1, "RandomBitsNode": 1, "ArgmaxNode": 2, - "ItemIntNode": 1, # temperature>0 condition only (no seed) + "ItemIntNode": 2, # top_k + temperature>0 condition only (no seed) "SoftmaxNode": 1, # top-p nucleus chain (top_p defaults to 1.0) } @@ -7747,12 +7749,14 @@ class SampleTopPTest(OpTestCase): "IfNode": 1, "RandomBitsNode": 1, "ArgmaxNode": 2, - "ItemIntNode": 2, + "ItemIntNode": 3, "SoftmaxNode": 1, - "SortNode": 1, + "SortNode": 2, "CumsumNode": 1, "MinNode": 1, - "WhereNode": 2, + "TakeNode": 1, + "ExpandDimsNode": 1, + "WhereNode": 3, } def create_model(self) -> nn.Module: @@ -7769,7 +7773,7 @@ def create_inputs(self) -> Tuple[torch.Tensor, ...]: @register_test class SampleTopKTest(OpTestCase): - """Top-k sample emits the extra threshold and combined mask nodes.""" + """Top-k sample emits the threshold before the top-p nucleus chain.""" name = "sample_top_k" skip_comparison = True # sampling RNG is not host/device bit-identical @@ -7779,13 +7783,13 @@ class SampleTopKTest(OpTestCase): "ArgmaxNode": 2, "ItemIntNode": 3, # seed + top_k + temperature>0 condition "SoftmaxNode": 1, - "SortNode": 1, + "SortNode": 2, "CumsumNode": 1, "MinNode": 1, "TakeNode": 1, "ExpandDimsNode": 1, - "LogicalOrNode": 1, - "WhereNode": 2, + "LogicalOrNode": 0, + "WhereNode": 3, } def create_model(self) -> nn.Module: diff --git a/backends/mlx/test/test_sample.py b/backends/mlx/test/test_sample.py index 276677d9993..5af803a23bc 100644 --- a/backends/mlx/test/test_sample.py +++ b/backends/mlx/test/test_sample.py @@ -97,8 +97,11 @@ def _sample( t = torch.tensor(float(temperature)) s = None if seed is None else torch.tensor(int(seed), dtype=torch.int64) p = torch.tensor(float(top_p)) # 1.0 = off - k = None if top_k is None else torch.tensor(int(top_k), dtype=torch.int64) - return torch.ops.mlx.sample(logits, t, p, s, k) + k = torch.tensor( + torch.iinfo(torch.int64).max if top_k is None else int(top_k), + dtype=torch.int64, + ) + return torch.ops.mlx.sample(logits, t, k, p, s) class TestSampleOp(unittest.TestCase): @@ -161,23 +164,31 @@ def test_top_p_one_keeps_all(self): self.assertTrue((tokens == 3).any()) def test_top_k_restricts_to_top_k(self): - # probs [0.5, 0.3, 0.15, 0.05]; top_k=2 keeps {0,1}. - base = torch.log(torch.tensor([0.5, 0.3, 0.15, 0.05])) + # Non-sorted probs [0.15, 0.5, 0.05, 0.3]; top_k=2 keeps {1,3}. + base = torch.log(torch.tensor([0.15, 0.5, 0.05, 0.3])) tokens = _sample(base.expand(5000, 4), 1.0, seed=0, top_k=2) - self.assertTrue(torch.isin(tokens, torch.tensor([0, 1])).all()) - self.assertEqual(set(tokens.tolist()), {0, 1}) + self.assertTrue(torch.isin(tokens, torch.tensor([1, 3])).all()) + self.assertEqual(set(tokens.tolist()), {1, 3}) - def test_top_k_none_keeps_all(self): + def test_top_k_default_keeps_all(self): # top_k=None -> no filtering; the tail token (index 3) is reachable. base = torch.log(torch.tensor([0.5, 0.3, 0.15, 0.05])) tokens = _sample(base.expand(20000, 4), 1.0, seed=0, top_k=None) self.assertTrue((tokens == 3).any()) + def test_top_k_clips_to_vocab_size(self): + # top_k > vocab is clipped to vocab size, so every token is reachable. + base = torch.log(torch.tensor([0.5, 0.3, 0.15, 0.05])) + tokens = _sample(base.expand(20000, 4), 1.0, seed=0, top_k=999) + self.assertEqual(set(tokens.tolist()), {0, 1, 2, 3}) + def test_top_k_and_top_p_compose(self): - # top_p=0.7 keeps {0,1}; top_k=1 intersects that to {0}. + # top_k is applied before top_p, so top_p sees renormalized top-k probs. + # top_k=3 -> [0.526, 0.316, 0.158]; top_p=0.83 keeps {0,1}. base = torch.log(torch.tensor([0.5, 0.3, 0.15, 0.05])) - tokens = _sample(base.expand(5000, 4), 1.0, seed=0, top_p=0.7, top_k=1) - self.assertEqual(set(tokens.tolist()), {0}) + tokens = _sample(base.expand(5000, 4), 1.0, seed=0, top_p=0.83, top_k=3) + self.assertTrue(torch.isin(tokens, torch.tensor([0, 1])).all()) + self.assertEqual(set(tokens.tolist()), {0, 1}) class TestSampleExport(unittest.TestCase): From 7e9ae359ff461dc54d3d213728122734d1f1fd44 Mon Sep 17 00:00:00 2001 From: Scott Roy Date: Mon, 29 Jun 2026 11:02:05 -0700 Subject: [PATCH 3/3] lint --- backends/mlx/ops.py | 4 +--- backends/mlx/test/test_ops.py | 6 +++--- 2 files changed, 4 insertions(+), 6 deletions(-) diff --git a/backends/mlx/ops.py b/backends/mlx/ops.py index 986b713a5ff..6a5be724f94 100644 --- a/backends/mlx/ops.py +++ b/backends/mlx/ops.py @@ -3628,9 +3628,7 @@ def emit_sample(): ) _, top_k_val = P.make_tmp_value_slot() P.emit( - ItemIntNode( - x=P.slot_to_tid(clipped_top_k), out=P.slot_to_vid(top_k_val) - ) + ItemIntNode(x=P.slot_to_tid(clipped_top_k), out=P.slot_to_vid(top_k_val)) ) _, top_k_index = P.make_tmp_value_slot() P.emit( diff --git a/backends/mlx/test/test_ops.py b/backends/mlx/test/test_ops.py index bd3a167f56a..8d9fcacde91 100644 --- a/backends/mlx/test/test_ops.py +++ b/backends/mlx/test/test_ops.py @@ -7702,7 +7702,7 @@ class SampleSeededTest(OpTestCase): "SortNode": 2, # top-k threshold + top-p nucleus chain "CumsumNode": 1, "MinNode": 1, - "TakeNode": 1, + "TakeNode": 2, # last-token slice + top-k threshold gather "ExpandDimsNode": 1, "WhereNode": 3, } @@ -7754,7 +7754,7 @@ class SampleTopPTest(OpTestCase): "SortNode": 2, "CumsumNode": 1, "MinNode": 1, - "TakeNode": 1, + "TakeNode": 2, # last-token slice + top-k threshold gather "ExpandDimsNode": 1, "WhereNode": 3, } @@ -7786,7 +7786,7 @@ class SampleTopKTest(OpTestCase): "SortNode": 2, "CumsumNode": 1, "MinNode": 1, - "TakeNode": 1, + "TakeNode": 2, # last-token slice + top-k threshold gather "ExpandDimsNode": 1, "LogicalOrNode": 0, "WhereNode": 3,