Skip to content

[ROCm] Fix EAGLE spec-decode verify silently sampling greedy on HIP - #37134

Merged
HaiShaw merged 10 commits into
sgl-project:mainfrom
xiaobochen-amd:rocm/eagle-verify-greedy-hip
Sep 16, 2026
Merged

HaiShaw merged 10 commits into
sgl-project:mainfrom
xiaobochen-amd:rocm/eagle-verify-greedy-hip

Conversation

@xiaobochen-amd

@xiaobochen-amd xiaobochen-amd commented Aug 30, 2026 •

Copy link
Copy Markdown
Contributor

Problem

EAGLE spec-decode verify has been committing argmax on ROCm, ignoring
temperature and top_p entirely, with nothing in the log saying so. At
temp > 0 it degenerates into repetition loops.

The cause is one line in eagle_sample:

if sampling_info.is_all_greedy or _is_cpu or _is_npu or _is_hip or _is_xpu:
    target_predict = torch.argmax(next_token_logits, dim=-1)

_is_hip sits next to is_all_greedy, so HIP takes the greedy path
unconditionally. It is there because the sampling branch needs
tree_speculative_sampling_target_only, top_k_renorm_prob and
top_p_renorm_prob, all of which are CUDA/MUSA-only in sgl_kernel.

Fix

Route HIP through the pure-Triton chain sampler in reject_sampling.py when
rejection 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 platform
it reduces to the predicate it replaces: with is_hip False the added term drops
out and the expression is the original one.

_renorm_top_k_top_p_hip follows the same shape as the op it replaces — a
threshold search, matching top_p_renorm_prob's ternary pivot search on
f(x) = sum(probs[probs > x]) — so it agrees with CUDA rather than with a
sort-based nucleus. top_p >= 1.0 short-circuits to an exact pass-through, to
stay 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 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) 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_fn dispatch. It used the eager A if cond else B form, which
resolves tree_speculative_sampling_target_only whether or not that branch is
taken. HIP must not import it at all, so the dispatch becomes an if/else
that imports inside the arm needing it. Non-HIP behaviour is unchanged.

One cfg read becomes resolved_view. _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 read becomes
resolved_view(server_args), which is what the declare_resolution docstring
prescribes 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 numerical
    equivalence against an independent torch reference, plus the top_p == 1.0
    exact-no-op and top_p -> 0 edges. Registered for both CUDA and AMD runners;
    the kernel is device-agnostic. 5 tests.

ruff, isort and black clean.


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.

@1am9trash 1am9trash left a comment •

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Sorry, I mis-clicked approve. Please check the following comments.

@1am9trash 1am9trash left a comment

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

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?

@xiaobochen-amd

Copy link
Copy Markdown
Contributor Author

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 dflash_utils.py does) puts the
renorm back on one code path. Private kernel, _is_hip branch, and its test gone;
kernels/ops/moe/test_renorm.py covers what's left.

Also collapsed the gate predicate to is_hip and not use_rejection_sampling.

3 files, +171/−15, down from 4 and +409.

@1am9trash 1am9trash left a comment

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Comments resolved.
Thanks for the update.

@jiejingzhangamd
jiejingzhangamd force-pushed the rocm/eagle-verify-greedy-hip branch from 5c6e07c to acd4d47 Compare September 8, 2026 18:10
@1am9trash

Copy link
Copy Markdown
Collaborator

Because of the new rebase commit pushed, the previous CI wae canceled, now I retriggered it (PR-test-base).

@jiejingzhangamd

Copy link
Copy Markdown
Contributor

Pushed a follow-up: HIP no longer auto-enables --speculative-use-rejection-sampling for EAGLE3.

That was crashing AMD stage-a-test-1-gpu-small-amd / test_basic_sanity_eagle3.py (TestBasicSanityEagle3.setUpClass) with draft vocab 32000 vs target 128256. Same-vocab EAGLE on HIP is unchanged.

@JohnQinAMD

Copy link
Copy Markdown
Contributor

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.

commit without it
carry real top_k into the draft proposal the draft can't tell a greedy request from a T=1 one, so it samples where it should take argmax
copy draft_probs out of the CUDA graph memory pool the pool recycles that block before eagle_sample reads it; 2-4% of tokens come back as an untrained padding row, GPQA drops to 0.4646
reject non-finite draft probabilities in the sampler the residual passes guard q with q == q, which infinities satisfy
a malformed draft probability must reject, not accept the accept test has no guard, so a corrupt q accepts unconditionally and commits the draft head's token

GPQA-Diamond, GLM-5.2-FP8, 8x MI355X, TP8, NEXTN steps=1 draft=2, 198q, T=1.0 top_p=0.95 max_tokens=98304:

base (bb15be6d79) this branch
exact_match 0.7071 ± 0.0324 0.8889 ± 0.0224
requests reaching a final answer 62.1% 94.4%
accuracy on requests that do finish 0.9268 0.9358
wall clock ~60 min ~24 min

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:

base this branch speculation off
T=0 0.941 0.940 0.951
T=1.0 top_p=0.95 0.928 0.925 0.937
accept length T=0 → T=1.0 1.932 → 1.931 1.931 → 1.886 —

No regression at either temperature — one question apart at T=0, four at T=1.0, against a 0.6pp sampling interval. The accept-length row is the one that matters: on base it does not move with temperature, because verify commits argmax whatever you ask for. This branch responds to it.

Performance

Same GSM8K sweep as above — 1319 questions, concurrency 16, so accuracy and throughput come from one run.

base this branch speculation off
output tok/s, T=0 746.3 744.1 601.4
output tok/s, T=1.0 747.3 727.1 577.3
end-to-end latency, T=0 181.4 s 180.5 s 222.3 s
end-to-end latency, T=1.0 181.7 s 191.5 s 233.4 s

T=0 is free (-0.3%, inside noise): is_all_greedy sends verify down the argmax branch either way. T=1.0 costs -2.7%, from acceptance dropping 1.931 → 1.886 as rejection sampling starts rejecting. A single request at concurrency 1 shows the same thing: 125.98 vs 128.76 tok/s at T=1.0, 135.28 vs 135.66 at T=0.

Against that, speculative decoding is worth +24% at T=0 and +26% at T=1.0 over no speculation at all, and on the GPQA run above this branch finishes 198 questions in ~24 minutes where base takes ~60 — the runaways it prevents cost far more compute than the rejections it adds.

CI: stage-a-test-1-gpu-small-amd now passes. The remaining red is a canceled runner task on base-b-test-1-gpu-small (5) that fast-failed every other CUDA shard, plus a missing libavutil.so.60 on XPU.

@JohnQinAMD
JohnQinAMD force-pushed the rocm/eagle-verify-greedy-hip branch from 9eed096 to 2cd1fcb Compare September 13, 2026 22:02
@xiaobochen-amd

Copy link
Copy Markdown
Contributor Author

/rerun-failed-ci

@At1a8

At1a8 commented Sep 14, 2026

Copy link
Copy Markdown
Contributor

Hi @xiaobochen-amd, could you please test the benchmark listed in #39253?

@xiaobochen-amd

Copy link
Copy Markdown
Contributor Author

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
0.7071 → 0.8889, requests reaching a final answer 62.1% → 94.4%, wall clock
60 → 24 min. Different model, but it looks like the same thing your Simple-QA
30.6% → 0.9% is measuring. GSM8K 1319 shows no regression.

@xiaobochen-amd

Copy link
Copy Markdown
Contributor Author

Re-ran the failing tests on MI355X in the same image CI uses
(rocm/sgl-dev:v0.5.19-rocm10-mi35x-20260913), against clean upstream/main:

main this PR
test_deterministic 2 failed / 4 error 2 failed / 4 error
test_sampling_mask 8 passed / 17 error 8 passed / 17 error
test_lean_attention 11 passed 11 passed
test_dsv4_indexer_quant 6 passed 6 passed
test_triton_mla_prefill_gfx950 3 passed 3 passed
test_triton_dense_prefill_gfx950 6 passed 6 passed
test_disaggregation_pp 3 passed 3 passed

Identical on every file, so none of them come from this branch.
base-b-test-1-gpu-large (1) is test_inkling_unified, a CUDA KL check that runs
no speculative decoding, so eagle_sample is never called there either.

@xiaobochen-amd

Copy link
Copy Markdown
Contributor Author

/rerun-failed-ci

JohnQinAMD and others added 10 commits September 15, 2026 06:29
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)
@At1a8

At1a8 commented Sep 16, 2026 •

Copy link
Copy Markdown
Contributor

Confirmed that this PR also fixes the accuracy issue(#39253) in DeepSeek-V4 MTP. Thanks @HaiShaw @xiaobochen-amd

dataset offcial(https://huggingface.co/deepseek-ai/DeepSeek-V4-Pro) this PR
GSM8K 92.6 96.44
AA-LCR 66.3 76.00
GPQA-Diamond 90.1 89.39
LongBench v2 51.5 53.28
Simple-QA 57.9 57.33
LiveCodeBench 93.5 90.43
MMLU-Pro 87.5 87.35
HLE 37.7 35.35

@HaiShaw
HaiShaw merged commit 7eedd57 into sgl-project:main Sep 16, 2026
372 of 438 checks passed
Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Labels

amd jit-kernel run-ci CI: run the baseline test suite on this PR speculative-decoding

Projects

None yet

Development

Successfully merging this pull request may close these issues.

6 participants