Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
16 changes: 14 additions & 2 deletions backends/mlx/custom_ops.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand All @@ -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)
Expand All @@ -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)
12 changes: 8 additions & 4 deletions backends/mlx/llm/sampling.py
Original file line number Diff line number Diff line change
Expand Up @@ -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.
Expand All @@ -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)
83 changes: 78 additions & 5 deletions backends/mlx/ops.py
Original file line number Diff line number Diff line change
Expand Up @@ -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)
Expand Down Expand Up @@ -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))
Expand Down Expand Up @@ -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(
Expand Down
62 changes: 55 additions & 7 deletions backends/mlx/test/test_ops.py
Original file line number Diff line number Diff line change
Expand Up @@ -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."""
Expand All @@ -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:
Expand All @@ -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)
}

Expand All @@ -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:
Expand All @@ -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
Expand Down
70 changes: 68 additions & 2 deletions backends/mlx/test/test_sample.py
Original file line number Diff line number Diff line change
Expand Up @@ -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)
Expand All @@ -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):
Expand Down Expand Up @@ -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
Expand Down Expand Up @@ -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()
Loading