Repository navigation
fix(sampling): treat top_k above the vocabulary as no top-k in the chain - #2254
Merged
Merged
Conversation
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
3 tasks done
This file contains hidden or bidirectional Unicode text that may be interpreted or compiled differently than what appears below. To review, open the file in an editor that reveals hidden Unicode characters.
Learn more about bidirectional Unicode characters
Sign up for free
to join this conversation on GitHub.
Already have an account?
Sign in to comment
Add this suggestion to a batch that can be applied as a single commit.This suggestion is invalid because no changes were made to the code.Suggestions cannot be applied while the pull request is closed.Suggestions cannot be applied while viewing a subset of changes.Only one suggestion per line can be applied in a batch.Add this suggestion to a batch that can be applied as a single commit.Applying suggestions on deleted lines is not supported.You must change the existing code in this line in order to create a valid suggestion.Outdated suggestions cannot be applied.This suggestion has been applied or marked resolved.Suggestions cannot be applied from pending reviews.Suggestions cannot be applied on multi-line comments.Suggestions cannot be applied while the pull request is queued to merge.Suggestion cannot be applied right now. Please check back later.
Summary
One chat completion with a
top_kabove the vocabulary abortedmlxcel-server(SIGABRT, exit 134): the stock sampling chain calledargpartition(-x, top_k - 1)unguarded, and the throw crossed a non-Resultcxx bridge function. Reproduction details are in the issue comment.fused_sample_filter_logitsnow treatstop_kat or above the vocabulary as inactive, the rulesampling_rejection_routesand the rejection kernel already use (llama.cpp clamps the same way). This is the single chain behindfused_sample,fused_sample_xtc,fused_sample_categorical,fused_sample_probsandfused_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: negativetop_kis already inactive, XTC special-token ids are range-checked, and top-logprobskis clamped in Rust.top_k_filteris unchanged (its callers guard); its doc now states thek < vocabprecondition.docs/python-client.mdstates thattop_k >= vocabkeeps every token.Tests
sampling_top_k_vocab_tests(mlxcel-core): every fused entry point, with and without XTC and top-p, attop_k64 (= vocab), 65, 1000000 andi32::MAXon a 64-token row. Each reported distribution equals thetop_k = 0one within 1e-6 and each draw is a valid id.tests/sampling_top_k_above_vocab_e2e.rs: bootsmlxcel-serveron Qwen3-0.6B-4bit and sendstop_k: 1000000withtop_p1.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):
top_p1.0, 0.9, omitted, and 0.9 withMLXCEL_SAMPLING_REJECTION=0all return HTTP 200 with the server alive;mlxcel generate --top-k 1000000with and without--top-p 0.9completes.cargo test -p mlxcel-core --lib sampling -- --test-threads=1: 203 passed. The e2e test passes.make verify-rocm-heldon 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 <= vocabis unchanged.Closes #2247