Skip to content

fix(sampling): treat top_k above the vocabulary as no top-k in the chain - #2254

Merged
inureyes merged 1 commit into
mainfrom
fix/issue-2247-top-k-above-vocab
Oct 8, 2026
Merged

inureyes merged 1 commit into
mainfrom
fix/issue-2247-top-k-above-vocab

Conversation

@inureyes

@inureyes inureyes commented Oct 8, 2026

Copy link
Copy Markdown
Member

Summary

One chat completion with a top_k above the vocabulary aborted mlxcel-server (SIGABRT, exit 134): the stock sampling chain called argpartition(-x, top_k - 1) unguarded, and the throw crossed a non-Result cxx bridge function. Reproduction details are in the issue comment.

  • fused_sample_filter_logits now treats top_k at or above the vocabulary as inactive, the rule sampling_rejection_routes and the rejection kernel already use (llama.cpp clamps the same way). This is the single chain behind fused_sample, fused_sample_xtc, fused_sample_categorical, fused_sample_probs and fused_sample_probs_xtc, so the server, mlxcel generate/run, the batch scheduler and speculative acceptance are all covered on every backend. Other request-derived inputs to the bridge were checked: negative top_k is already inactive, XTC special-token ids are range-checked, and top-logprobs k is clamped in Rust.
  • The Rust reference top_k_filter is unchanged (its callers guard); its doc now states the k < vocab precondition.
  • docs/python-client.md states that top_k >= vocab keeps every token.

Tests

  • sampling_top_k_vocab_tests (mlxcel-core): every fused entry point, with and without XTC and top-p, at top_k 64 (= vocab), 65, 1000000 and i32::MAX on a 64-token row. Each reported distribution equals the top_k = 0 one within 1e-6 and each draw is a valid id.
  • tests/sampling_top_k_above_vocab_e2e.rs: boots mlxcel-server on Qwen3-0.6B-4bit and sends top_k: 1000000 with top_p 1.0 and 0.9; asserts HTTP 200, sampled tokens, and a live, healthy server. Skips when the checkpoint is absent.

Before the fix (gfx1151): the unit test binary aborted ([argpartition] Received invalid kth 64 ... shape: (1,64), exit 134) and the e2e test failed with no response because the server aborted.

Verification

On gfx1151 (ROCm):

  • Server and CLI reproduction after the fix: top_p 1.0, 0.9, omitted, and 0.9 with MLXCEL_SAMPLING_REJECTION=0 all return HTTP 200 with the server alive; mlxcel generate --top-k 1000000 with and without --top-p 0.9 completes.
  • cargo test -p mlxcel-core --lib sampling -- --test-threads=1: 203 passed. The e2e test passes.
  • Fast gates (versions, kernel dtype keys, kernel port dispatch, llama-compat, fmt) and clippy on the changed targets pass.
  • make verify-rocm-held on 8c39200 (base 7a3fcc4, current main): OK, 12197 passed, 0 failed, 403 ignored across 164 test binaries; smoke and workspace clippy pass.

Not verified: Metal and CUDA (not available on this host). The change is in shared bridge code they also run; it only removes a top-k filter that kept every token, so their output for top_k <= vocab is unchanged.

Closes #2247

The stock sampling chain called argpartition(-x, top_k - 1) without comparing top_k to the vocabulary. argpartition throws for top_k > vocab, and the fused sampler bridge functions are not Result, so the throw aborted the process. The server forwards a request's top_k unchanged, so one chat completion with top_k 1000000 and top_p 1.0 (or any top_p with MLXCEL_SAMPLING_REJECTION=0) killed mlxcel-server with SIGABRT; mlxcel generate --top-k 1000000 aborted the same way.

fused_sample_filter_logits now treats top_k at or above the vocabulary as inactive, the rule sampling_rejection_routes and the rejection kernel already used. This covers fused_sample, fused_sample_xtc, fused_sample_categorical and both fused_sample_probs variants, on every backend, since they share the chain.

Tests: sampling_top_k_vocab_tests runs every entry point with top_k equal to, one above, and far above a 64-token vocabulary and compares the reported distribution to top_k = 0; tests/sampling_top_k_above_vocab_e2e.rs boots mlxcel-server on Qwen3-0.6B-4bit and sends a huge top_k. Both aborted or failed before the fix on gfx1151.

Closes #2247
@inureyes inureyes added status:done Completed type:bug Bug fixes, error corrections, or issue resolutions priority:high High priority area:inference Generation, sampling, decoding (incl. speculative, DRY) labels Oct 8, 2026
@inureyes
inureyes merged commit d078601 into main Oct 8, 2026
25 of 27 checks passed
@inureyes
inureyes deleted the fix/issue-2247-top-k-above-vocab branch October 8, 2026 23:52
Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Labels

area:inference Generation, sampling, decoding (incl. speculative, DRY) priority:high High priority status:done Completed type:bug Bug fixes, error corrections, or issue resolutions

Projects

None yet

Development

Successfully merging this pull request may close these issues.

fix(sampling): treat top_k above the vocabulary as no top-k in the C++ chain

1 participant