Repository navigation
[ROCm] Fix EAGLE spec-decode verify silently sampling greedy on HIP - #37134
Conversation
1am9trash
left a comment
There was a problem hiding this comment.
Hi, @xiaobochen-amd
python/sglang/kernels/ops/sampling/renorm_triton.py already provides rocm triton replacements for top_k_renorm_prob / top_p_renorm_prob, and srt/speculative/dflash_utils.py already consumes them via a platform-conditional import that aliases them to the same names:
elif is_hip():
from sglang.kernels.ops.sampling.renorm_triton import (
top_k_renorm_probs_triton as top_k_renorm_prob,
)
from sglang.kernels.ops.sampling.renorm_triton import (
top_p_renorm_probs_triton as top_p_renorm_prob,
)Should we reuse those instead of adding new implementations?
@1am9trash Done — same signatures, so the alias import (as Also collapsed the gate predicate to 3 files, +171/−15, down from 4 and +409. |
1am9trash
left a comment
There was a problem hiding this comment.
Comments resolved.
Thanks for the update.
5c6e07c to
acd4d47
Compare
|
Because of the new rebase commit pushed, the previous CI wae canceled, now I retriggered it (PR-test-base). |
|
Pushed a follow-up: HIP no longer auto-enables That was crashing AMD |
|
I pushed four commits onto this branch. They fix bugs that only become reachable because of this PR — before it HIP never entered the sampling-verify path, so the code was dead. Landing without them ships the corruption in row 2.
GPQA-Diamond, GLM-5.2-FP8, 8x MI355X, TP8, NEXTN
Accuracy on requests that finish is the same on both sides — the delta is entirely whether the model terminates. Without this PR, verify commits argmax whatever temperature was asked for, and greedy decoding on long chain-of-thought loops until the budget is gone. GSM8K, all 1319 questions, 5-shot:
No regression at either temperature — one question apart at PerformanceSame GSM8K sweep as above — 1319 questions, concurrency 16, so accuracy and throughput come from one run.
Against that, speculative decoding is worth +24% at CI: |
9eed096 to
2cd1fcb
Compare
|
/rerun-failed-ci |
|
Hi @xiaobochen-amd, could you please test the benchmark listed in #39253? |
@At1a8 We haven't run that set. @JohnQinAMD measured this branch on GLM-5.2-FP8 (NEXTN 1/2): GPQA-Diamond |
|
Re-ran the failing tests on MI355X in the same image CI uses
Identical on every file, so none of them come from this branch. |
|
/rerun-failed-ci |
The sgl_kernel sampling-verify ops (tree_speculative_sampling_target_only, top_k_renorm_prob, top_p_renorm_prob) are CUDA/MUSA-only, so HIP fell through eagle_sample's gate to argmax, silently ignoring temperature and top_p and degenerating into repetition loops at temp > 0. Route HIP through the pure-Triton chain sampler (reject_sampling.py) when rejection sampling is active and the batch is not all-greedy, with a Triton pivot-search renorm (_renorm_top_k_top_p_hip) standing in for the CUDA-only renorm ops. top_p >= 1.0 is an exact no-op there, to stay bit-exact with CUDA's top_p_renorm_prob(p, 1.0) == p. The gate is factored into a pure helper, _verify_uses_greedy, which reduces to the original predicate on every non-HIP platform: with is_hip False the added term drops out and the expression is the one it replaced. Two shapes here come from this branch rather than from the original fix, because the surrounding code has moved: - sampling_fn used the eager `A if cond else B` form, which resolves tree_speculative_sampling_target_only whether or not that branch is taken. Since HIP must not import it at all, the dispatch becomes an if/else that imports inside the arm that needs it. Non-HIP behaviour is unchanged. - _handle_eagle_family writes through declare_resolution, whose contract is that a field it decides is not visible to a later `cfg.` read. The rejection-sampling validation immediately below was reading cfg, so it would have skipped every auto-enabled config -- including its own "rejection sampling is enabled" log. That one read becomes resolved_view(server_args), which is what the declare_resolution docstring prescribes for exactly this case. Defaulting the flag on matters because without it the fix only reaches users who pass --speculative-use-rejection-sampling; everyone else still silently gets argmax verify. It flips only on HIP, for EAGLE/EAGLE3 with topk=1, default accept thresholds and non-deterministic inference -- the same conditions the validation below raises on, so the flip can never turn a working config into an error. Setting it also makes the draft worker emit the target-vocab proposal (draft_probs) that the Triton chain sampler consumes. Tests: a gate-routing truth table (CPU) and renorm-vs-nucleus numerical equivalence (GPU).
The blank line predates this PR, but pre-commit formats every file a change touches, so it has to go here.
check_registered_tests.py admits a newly added registered test only under test/registered/<kind>/<subsystem>/, and `spec` is not one of the kinds. This test drives a GPU and compares the Triton renorm against nucleus, so `kernel` is the kind that fits -- `unit` is restricted to CPU suites. Its suites move with it: base-b-kernel-unit on CUDA and jit-kernel-unit-test-amd on ROCm, both of which carry the *-kernel-* name that kind requires. test_eagle_gate_routing.py is a CPU truth table and already sat correctly under unit/spec.
…mentation renorm_triton.py already carries top_k_renorm_probs_triton and top_p_renorm_probs_triton, and dflash_utils.py already aliases them to the sgl_kernel names on HIP. Do the same here rather than keeping a private pivot kernel: the renorm block goes back to one code path for both platforms, and the ops that survive are the ones test/registered/kernels/ops/moe/test_renorm.py already covers. Drops _top_p_renorm_kernel, _renorm_top_k_top_p_hip, the _is_hip branch around the renorm calls, and the renorm equivalence test that only existed to guard the private implementation.
Co-authored-by: Cursor <cursoragent@cursor.com>
EAGLE3 hot-token vocabs cannot feed the rejection kernel yet, so defaulting the flag on crashed stage-a test_basic_sanity_eagle3. Leave that path greedy until the worker can scatter draft probs into the target vocab. Co-authored-by: Cursor <cursoragent@cursor.com>
sample_draft_proposal decides greedy-vs-sample from temperature alone, but
SamplingParams rewrites temperature 0 to temperature=1.0 with top_k=1, so a
greedy request is indistinguishable from a T=1 one there. It samples a sharp
but non-degenerate distribution and proposes a non-argmax token often enough
to break the draft chain, costing accept length.
Pass the per-request top_ks through and let a greedy row propose its argmax.
The CUDA graph runner has to carry top_ks in a device buffer the same way it
already carries temperatures: its synthetic SamplingBatchInfo used a host-side
placeholder, so without this the correction never sees a real top_k and only
about a fifth of the loss comes back.
TOP_K_ALL rather than -1 for the buffer fill: -1 is not a top_k this pipeline
ever carries, and it reads as top_k <= 1, i.e. greedy, for the padded rows.
GLM-5.2-MXFP4, MI355X, TP4/EP4, EAGLE steps=5 topk=1 draft=6,
GSM8K 200q 5-shot temp=0 conc=8:
baseline before after
accept len 3.864 3.117 3.857
output tok/s 653.3 569.0 653.1
GSM8K accuracy 0.929 - 0.936
Co-Authored-By: Claude Opus 5 (1M context) <noreply@anthropic.com>
(cherry picked from commit 72ea484)
draft_probs is the one draft-graph output that the target forward does not
consume: it is read afterwards, by eagle_sample. The tensor the replay
returns points into the graph's private memory pool, and that block is
recycled before the read. What lands there is the DSA top-k mask, so
eagle_sample sees -inf in draft_probs -- scattered runs whose position
tracks the sequence at roughly the accept length, each ending on a
64-aligned column, always at the same draft step because the pool's
allocation order is fixed at capture.
Downstream that turns into corrupt output: the residual p - q goes
infinite, norm_sum with it, the cumulative search never matches, and the
sampler commits its degenerate fallback, an untrained embedding row. The
model degenerates from there.
Copy out at the graph boundary. Only rejection sampling produces
draft_probs, so this is a no-op everywhere else.
GLM-5.2-MXFP4, MI355X, TP4/EP4, EAGLE steps=5 topk=1 draft=6,
rejection sampling on, temperature=1.0 top_p=0.95:
before after
non-finite entries in draft_probs 5 to 59 0
stray tokens per 1200 generated 25 to 81 0
The same three arms with CUDA graphs disabled show zero in both columns,
which is what identified the pool rather than the producer.
GSM8K 200q 5-shot at temperature=0: 0.940 / 0.000 invalid / accept len
3.035 / 512.8 tok/s, against 0.935 / 0.000 / 3.031 / 520.3 without the
copy. The 1.4% throughput gap is inside this harness's noise: three
repeats on one unchanged server span 582 to 646 tok/s.
Co-Authored-By: Claude Opus 5 (1M context) <noreply@anthropic.com>
(cherry picked from commit 3e1fae0)
Defence in depth behind the previous commit. The residual guard tested q for NaN with `q_val == q_val`, which infinities satisfy, so a malformed draft distribution reaches `p - q` and drives norm_sum infinite. The cumulative search then never matches and the kernel commits its fallback, VOCAB_SIZE - 1 -- an id past the tokenizer's vocabulary whenever the embedding is padded, as it is for GLM-5.2 (154880 against 154856). Test q against the probability range instead. A comparison with NaN is false, so the range test rejects NaN along with the infinities and with negatives. A sampling kernel should not be able to emit an invalid token id because an input was malformed, whatever wrote that input. Co-Authored-By: Claude Opus 5 (1M context) <noreply@anthropic.com> (cherry picked from commit dee891b)
) The previous commit guards q in the two residual passes but not in the accept test itself, where it does the most damage: `coin * q < p` is -inf < p for an -inf q and 0 < p for a zero one, so a corrupt row accepts unconditionally and the committed token comes from the draft head rather than the target. That is the direction that costs output quality, and it is invisible -- acceptance climbing to 1.0 reads as a good draft model. Zero also passes the range test the residual passes use, so widening that guard alone does not cover this site. X was sampled from q, so q(X) has to be strictly positive; anything else means the row is not the distribution X came from. Reject on that, which sends the step through the residual path and resamples from the target. Co-authored-by: Claude Opus 5 (1M context) <noreply@anthropic.com> (cherry picked from commit 8d63553)
262741a to
69cf150
Compare
|
Confirmed that this PR also fixes the accuracy issue(#39253) in DeepSeek-V4 MTP. Thanks @HaiShaw @xiaobochen-amd
|
Problem
EAGLE spec-decode verify has been committing
argmaxon ROCm, ignoringtemperatureandtop_pentirely, with nothing in the log saying so. Attemp > 0it degenerates into repetition loops.The cause is one line in
eagle_sample:_is_hipsits next tois_all_greedy, so HIP takes the greedy pathunconditionally. It is there because the sampling branch needs
tree_speculative_sampling_target_only,top_k_renorm_probandtop_p_renorm_prob, all of which are CUDA/MUSA-only insgl_kernel.Fix
Route HIP through the pure-Triton chain sampler in
reject_sampling.pywhenrejection sampling is active and the batch is not all-greedy, with a Triton
renorm (
_renorm_top_k_top_p_hip) standing in for the CUDA-only renorm ops.The gate becomes a pure helper,
_verify_uses_greedy. On every non-HIP platformit reduces to the predicate it replaces: with
is_hipFalse the added term dropsout and the expression is the original one.
_renorm_top_k_top_p_hipfollows the same shape as the op it replaces — athreshold search, matching
top_p_renorm_prob's ternary pivot search onf(x) = sum(probs[probs > x])— so it agrees with CUDA rather than with asort-based nucleus.
top_p >= 1.0short-circuits to an exact pass-through, tostay bit-exact with CUDA's
top_p_renorm_prob(p, 1.0) == p. top-k uses a sort,skipped entirely when no row restricts it.
Defaulting the flag on matters: without it the fix only reaches users who pass
--speculative-use-rejection-sampling, and everyone else still silently getsargmax verify. It flips only on HIP, for EAGLE/EAGLE3 with
topk=1, defaultaccept thresholds and non-deterministic inference — the same conditions the
validation below raises on, so the flip can never turn a working config into an
error. Setting it also makes the draft worker emit the target-vocab proposal
(
draft_probs) the Triton chain sampler consumes.Two shapes that come from the current tree, not from the fix
Worth pointing out so they don't read as gratuitous:
sampling_fndispatch. It used the eagerA if cond else Bform, whichresolves
tree_speculative_sampling_target_onlywhether or not that branch istaken. HIP must not import it at all, so the dispatch becomes an
if/elsethat imports inside the arm needing it. Non-HIP behaviour is unchanged.
One
cfgread becomesresolved_view._handle_eagle_familywrites throughdeclare_resolution, whose contract is that a field it decides is not visibleto a later
cfg.read. The rejection-sampling validation immediately below wasreading
cfg, so it would have skipped every auto-enabled config — including itsown "rejection sampling is enabled" log. That read becomes
resolved_view(server_args), which is what thedeclare_resolutiondocstringprescribes for exactly this case.
Test
test/registered/unit/spec/test_eagle_gate_routing.py— gate truth table, CPU.Pins that the helper reduces to the original predicate off HIP, and that HIP
only leaves greedy when rejection sampling is on and the batch is not
all-greedy. 4 tests.
test/registered/spec/eagle/test_rocm_eagle_renorm.py— renorm numericalequivalence against an independent torch reference, plus the
top_p == 1.0exact-no-op and
top_p -> 0edges. Registered for both CUDA and AMD runners;the kernel is device-agnostic. 5 tests.
ruff,isortandblackclean.CI States
Latest PR Test (Base): ✅ Run #34937231207
Latest PR Test (Extra): ❌ Run #34937230795
Latest PR Test (AMD ROCm 10): ➖ No AMD PR run found for this commit.