diff --git a/backends/mlx/custom_ops.py b/backends/mlx/custom_ops.py index 17c07097f70..5605b59c543 100644 --- a/backends/mlx/custom_ops.py +++ b/backends/mlx/custom_ops.py @@ -397,15 +397,19 @@ def gather_qmm_fake( def sample( logits: Tensor, temperature: Tensor, + top_k: Tensor, top_p: Tensor, seed: Optional[Tensor] = None, ) -> Tensor: """ - Gumbel-max sampling from softmax(logits / temperature), with top-p (nucleus). + 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 @@ -422,6 +426,14 @@ 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) cum = torch.cumsum(s_probs, dim=-1) @@ -440,5 +452,5 @@ def sample( @torch.library.register_fake("mlx::sample") -def sample_fake(logits, temperature, top_p, seed=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 a059e5cff08..fb94ceb7d27 100644 --- a/backends/mlx/llm/sampling.py +++ b/backends/mlx/llm/sampling.py @@ -19,7 +19,9 @@ 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 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. @@ -31,10 +33,12 @@ 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 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_k, top_p, seed) diff --git a/backends/mlx/ops.py b/backends/mlx/ops.py index 86e322a16e7..6a5be724f94 100644 --- a/backends/mlx/ops.py +++ b/backends/mlx/ops.py @@ -3528,10 +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, 4, "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 + 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) @@ -3612,8 +3612,82 @@ 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), + ) + ) + + _, 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; SortNode is ascending-only, so sort -probs for descending. + # 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)) @@ -3674,7 +3748,6 @@ def emit_sample(): 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 afd4f276dde..8d9fcacde91 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.""" @@ -7686,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": 2, # last-token slice + top-k threshold gather + "ExpandDimsNode": 1, + "WhereNode": 3, } def create_model(self) -> nn.Module: @@ -7715,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) } @@ -7736,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": 2, # last-token slice + top-k threshold gather + "ExpandDimsNode": 1, + "WhereNode": 3, } def create_model(self) -> nn.Module: @@ -7756,6 +7771,39 @@ def create_inputs(self) -> Tuple[torch.Tensor, ...]: ) +@register_test +class SampleTopKTest(OpTestCase): + """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 + expected_node_counts = { + "IfNode": 1, + "RandomBitsNode": 1, + "ArgmaxNode": 2, + "ItemIntNode": 3, # seed + top_k + temperature>0 condition + "SoftmaxNode": 1, + "SortNode": 2, + "CumsumNode": 1, + "MinNode": 1, + "TakeNode": 2, # last-token slice + top-k threshold gather + "ExpandDimsNode": 1, + "LogicalOrNode": 0, + "WhereNode": 3, + } + + 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..5af803a23bc 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,21 @@ 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 = 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): @@ -142,6 +163,33 @@ 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): + # 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([1, 3])).all()) + self.assertEqual(set(tokens.tolist()), {1, 3}) + + 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_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.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): """Runtime-input semantics that survive export: temperature and seed stay @@ -218,6 +266,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()