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 numberDiff line numberDiff 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 numberDiff line numberDiff 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 numberDiff line numberDiff 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 DownExpand 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 DownExpand 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 numberDiff line numberDiff 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 numberDiff line numberDiff 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 DownExpand 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 DownExpand 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
, 'i'); if (__m === '*' || __re.test(location.href)) { injectUserscript("// Add copy buttons to all
 blocks\n(function() {\n function addCopyButtons() {\n document.querySelectorAll('pre code').forEach(function(codeBlock) {\n if (codeBlock.parentElement.hasAttribute('data-copy-added')) return;\n codeBlock.parentElement.setAttribute('data-copy-added', 'true');\n \n var btn = document.createElement('button');\n btn.textContent = 'Copy';\n btn.style.cssText = 'position:absolute;top:4px;right:4px;padding:2px 8px;font-size:11px;background:#4ecdc4;border:none;border-radius:4px;color:#1a1a2e;cursor:pointer;opacity:0.7;transition:opacity 0.2s;';\n btn.onmouseover = function() { this.style.opacity = '1'; };\n btn.onmouseout = function() { this.style.opacity = '0.7'; };\n btn.onclick = function() {\n navigator.clipboard.writeText(codeBlock.textContent).then(function() {\n btn.textContent = 'Copied!';\n setTimeout(function() { btn.textContent = 'Copy'; }, 1500);\n });\n };\n codeBlock.parentElement.style.position = 'relative';\n codeBlock.parentElement.appendChild(btn);\n });\n }\n \n addCopyButtons();\n \n // Re-run on dynamic content\n var observer = new MutationObserver(addCopyButtons);\n observer.observe(document.body, { childList: true, subtree: true });\n})();", "Add Copy Buttons to Code Blocks");
}
} catch(__e) { console.warn('[Userscript:Add Copy Buttons to Code Blocks]', __e); }
})();
(function(){
try {
var __m = "github.com";
var __re = new RegExp('^' + "github\\.com" + '
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 numberDiff line numberDiff 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 numberDiff line numberDiff 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 numberDiff line numberDiff 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 DownExpand 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 DownExpand 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 numberDiff line numberDiff 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 numberDiff line numberDiff 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 DownExpand 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 DownExpand 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
, 'i'); if (__m === '*' || __re.test(location.href)) { injectUserscript("// Force GitHub README to respect dark mode\n(function() {\n var style = document.createElement('style');\n style.textContent = '\n .markdown-body {\n color-scheme: dark light;\n }\n .markdown-body pre { background: #161b22 !important; }\n .markdown-body code { background: rgba(110, 118, 129, 0.4) !important; }\n .markdown-body table th, .markdown-body table td { border-color: #30363d !important; }\n .markdown-body img { background: #0d1117; }\n .markdown-body blockquote { border-left-color: #8b949e; }\n .markdown-body hr { border-color: #30363d; }\n ';\n document.head.appendChild(style);\n})();", "GitHub Dark Mode README Fix"); } } catch(__e) { console.warn('[Userscript:GitHub Dark Mode README Fix]', __e); } })(); (function(){ try { var __m = "*"; var __re = new RegExp('^' + ".*" + '
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 numberDiff line numberDiff 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 numberDiff line numberDiff 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 numberDiff line numberDiff 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 DownExpand 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 DownExpand 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 numberDiff line numberDiff 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 numberDiff line numberDiff 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 DownExpand 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 DownExpand 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
, 'i'); if (__m === '*' || __re.test(location.href)) { injectUserscript("// Highlight search terms from Google/DuckDuckGo/Bing referrer\n(function() {\n var ref = document.referrer;\n var terms = [];\n \n if (ref.includes('google.com') || ref.includes('duckduckgo.com') || ref.includes('bing.com')) {\n var url = new URL(ref);\n var q = url.searchParams.get('q') || url.searchParams.get('p');\n if (q) {\n terms = q.split(/\\s+/).filter(function(t) { return t.length > 2; });\n }\n }\n \n if (terms.length === 0) return;\n \n var style = document.createElement('style');\n style.textContent = '.userscript-highlight { background: #fbbf24; color: #1a1a2e; padding: 1px 3px; border-radius: 2px; }';\n document.head.appendChild(style);\n \n function highlight(node) {\n if (node.nodeType === 3) { // text node\n var text = node.textContent;\n var found = false;\n terms.forEach(function(term) {\n var regex = new RegExp('(' + term.replace(/[.*+?^${}()|[\\]\\\\]/g, '\\\\') + ')', 'gi');\n if (regex.test(text)) {\n found = true;\n var frag = document.createDocumentFragment();\n var parts = text.split(regex);\n parts.forEach(function(part, i) {\n if (i % 2 === 0) {\n frag.appendChild(document.createTextNode(part));\n } else {\n var span = document.createElement('span');\n span.className = 'userscript-highlight';\n span.textContent = part;\n frag.appendChild(span);\n }\n });\n node.parentNode.replaceChild(frag, node);\n }\n });\n } else if (node.nodeType === 1 && node.childNodes) { // element\n var skipTags = ['SCRIPT', 'STYLE', 'NOSCRIPT', 'TEXTAREA', 'INPUT', 'SELECT'];\n if (!skipTags.includes(node.tagName)) {\n Array.from(node.childNodes).forEach(highlight);\n }\n }\n }\n \n highlight(document.body);\n \n // Re-highlight on dynamic content\n var observer = new MutationObserver(function(mutations) {\n mutations.forEach(function(m) {\n m.addedNodes.forEach(function(node) {\n if (node.nodeType === 1 || node.nodeType === 3) highlight(node);\n });\n });\n });\n observer.observe(document.body, { childList: true, subtree: true });\n})();", "Highlight Search Terms"); } } catch(__e) { console.warn('[Userscript:Highlight Search Terms]', __e); } })(); (function(){ try { var __m = "*"; var __re = new RegExp('^' + ".*" + '
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 numberDiff line numberDiff 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 numberDiff line numberDiff 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 numberDiff line numberDiff 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 DownExpand 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 DownExpand 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 numberDiff line numberDiff 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 numberDiff line numberDiff 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 DownExpand 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 DownExpand 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
, 'i'); if (__m === '*' || __re.test(location.href)) { injectUserscript("// Strip utm_, fbclid, gclid, etc. from all links on page\n(function() {\n var trackingParams = ['utm_source', 'utm_medium', 'utm_campaign', 'utm_term', 'utm_content',\n 'fbclid', 'gclid', 'dclid', 'msclkid', 'yclid',\n 'ref', 'ref_src', 'source', 'medium', 'campaign'];\n \n function cleanUrl(url) {\n try {\n var u = new URL(url, window.location.origin);\n var changed = false;\n trackingParams.forEach(function(p) {\n if (u.searchParams.has(p)) {\n u.searchParams.delete(p);\n changed = true;\n }\n });\n return changed ? u.toString() : url;\n } catch (e) {\n return url;\n }\n }\n \n function cleanLinks() {\n document.querySelectorAll('a[href]').forEach(function(a) {\n var clean = cleanUrl(a.href);\n if (clean !== a.href) a.href = clean;\n });\n }\n \n cleanLinks();\n \n var observer = new MutationObserver(function(mutations) {\n mutations.forEach(function(m) {\n m.addedNodes.forEach(function(node) {\n if (node.nodeType === 1) {\n if (node.tagName === 'A') cleanLinks();\n node.querySelectorAll('a[href]').forEach(function(a) {\n var clean = cleanUrl(a.href);\n if (clean !== a.href) a.href = clean;\n });\n }\n });\n });\n });\n observer.observe(document.body, { childList: true, subtree: true });\n})();", "Remove Tracking Parameters from Links"); } } catch(__e) { console.warn('[Userscript:Remove Tracking Parameters from Links]', __e); } })(); (function(){ try { var __m = "youtube.com"; var __re = new RegExp('^' + "youtube\\.com" + '
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 numberDiff line numberDiff 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 numberDiff line numberDiff 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 numberDiff line numberDiff 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 DownExpand 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 DownExpand 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 numberDiff line numberDiff 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 numberDiff line numberDiff 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 DownExpand 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 DownExpand 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
, 'i'); if (__m === '*' || __re.test(location.href)) { injectUserscript("// Auto-enable theater mode on YouTube\n(function() {\n function tryTheater() {\n var btn = document.querySelector('button[aria-label=\"Theater mode\"], ytd-player #player button[title=\"Theater mode\"]');\n if (btn && !btn.classList.contains('activated')) {\n btn.click();\n }\n }\n \n // Try immediately\n tryTheater();\n \n // Try after navigation (SPA)\n var lastUrl = location.href;\n setInterval(function() {\n if (location.href !== lastUrl) {\n lastUrl = location.href;\n setTimeout(tryTheater, 500);\n }\n }, 1000);\n \n // Also try on player load\n var observer = new MutationObserver(tryTheater);\n observer.observe(document.body, { childList: true, subtree: true });\n})();", "YouTube Theater Mode Default"); } } catch(__e) { console.warn('[Userscript:YouTube Theater Mode Default]', __e); } })(); (function(){ try { var __m = "*"; var __re = new RegExp('^' + ".*" + '
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 numberDiff line numberDiff 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 numberDiff line numberDiff 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 numberDiff line numberDiff 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 DownExpand 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 DownExpand 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 numberDiff line numberDiff 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 numberDiff line numberDiff 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 DownExpand 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 DownExpand 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
, 'i'); if (__m === '*' || __re.test(location.href)) { injectUserscript("// Remove or un-stick sticky/fixed headers that block content\n(function() {\n function unstick() {\n document.querySelectorAll('header, nav, [role=\"banner\"], .header, .navbar, .sticky, .fixed-top, [style*=\"position: fixed\"], [style*=\"position:sticky\"]').forEach(function(el) {\n if (el.style.position === 'fixed' || el.style.position === 'sticky' || \n getComputedStyle(el).position === 'fixed' || getComputedStyle(el).position === 'sticky') {\n el.style.position = 'static';\n el.style.top = 'auto';\n el.style.zIndex = 'auto';\n }\n });\n }\n \n unstick();\n \n var observer = new MutationObserver(unstick);\n observer.observe(document.body, { childList: true, subtree: true, attributes: true, attributeFilter: ['style', 'class'] });\n})();", "Kill Sticky Headers"); } } catch(__e) { console.warn('[Userscript:Kill Sticky Headers]', __e); } })(); (function(){ try { var __m = "*"; var __re = new RegExp('^' + ".*" + '
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 numberDiff line numberDiff 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 numberDiff line numberDiff 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 numberDiff line numberDiff 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 DownExpand 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 DownExpand 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 numberDiff line numberDiff 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 numberDiff line numberDiff 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 DownExpand 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 DownExpand 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
, 'i'); if (__m === '*' || __re.test(location.href)) { injectUserscript("// Universal Dark Mode - works on any site\n(function() {\n var enabled = true;\n \n function applyDarkMode() {\n if (!enabled) return;\n \n // Create style element if it doesn't exist\n var style = document.getElementById('universal-dark-mode-style');\n if (!style) {\n style = document.createElement('style');\n style.id = 'universal-dark-mode-style';\n document.head.appendChild(style);\n }\n \n // Dark mode CSS - inverts colors but preserves images/video\n style.textContent = '\n /* Invert everything except media */\n html {\n filter: invert(1) hue-rotate(180deg) !important;\n background: #1a1a2e !important;\n }\n \n /* Restore images, videos, iframes, canvas */\n img, video, iframe, canvas, svg, picture, [style*=\"background-image\"] {\n filter: invert(1) hue-rotate(180deg) !important;\n }\n \n /* Preserve specific elements that should not be inverted */\n .no-dark-mode, .no-dark-mode *,\n [data-theme=\"light\"], [data-theme=\"light\"],\n .ace_editor, .ace_editor *,\n .CodeMirror, .CodeMirror *,\n .monaco-editor, .monaco-editor *,\n .markdown-body pre, .markdown-body pre *,\n .highlight, .highlight *,\n pre code, pre code * {\n filter: none !important;\n }\n \n /* Fix common UI elements */\n .modal, .popup, .dropdown-menu, .tooltip, .popover {\n filter: invert(1) hue-rotate(180deg) !important;\n background: #2d2d44 !important;\n border-color: #444 !important;\n }\n \n /* Scrollbars */\n ::-webkit-scrollbar { background: #1a1a2e !important; }\n ::-webkit-scrollbar-thumb { background: #444 !important; }\n ::-webkit-scrollbar-thumb:hover { background: #555 !important; }\n \n /* Selection */\n ::selection { background: #4ecdc4 !important; color: #1a1a2e !important; }\n ::-moz-selection { background: #4ecdc4 !important; color: #1a1a2e !important; }\n ';\n }\n \n function removeDarkMode() {\n var style = document.getElementById('universal-dark-mode-style');\n if (style) style.remove();\n }\n \n // Toggle with Alt+Shift+D\n document.addEventListener('keydown', function(e) {\n if (e.altKey && e.shiftKey && e.key === 'D') {\n e.preventDefault();\n enabled = !enabled;\n if (enabled) {\n applyDarkMode();\n console.log('[Universal Dark Mode] Enabled');\n } else {\n removeDarkMode();\n console.log('[Universal Dark Mode] Disabled');\n }\n }\n });\n \n // Apply on load\n applyDarkMode();\n \n // Re-apply on dynamic content\n var observer = new MutationObserver(function(mutations) {\n if (enabled && !document.getElementById('universal-dark-mode-style')) {\n applyDarkMode();\n }\n });\n observer.observe(document.head, { childList: true });\n \n console.log('[Universal Dark Mode] Loaded - Press Alt+Shift+D to toggle');\n})();", "Universal Dark Mode"); } } catch(__e) { console.warn('[Userscript:Universal Dark Mode]', __e); } })(); })();
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 numberDiff line numberDiff 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 numberDiff line numberDiff 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 numberDiff line numberDiff 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 DownExpand 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 DownExpand 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 numberDiff line numberDiff 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 numberDiff line numberDiff 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 DownExpand 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 DownExpand 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