Skip to content

[ROCm] EAGLE: carry real top_k into the draft proposal - #55

Merged
JohnQinAMD merged 1 commit into
mainfrom
fix/eagle-draft-greedy-topk
Sep 11, 2026
Merged

JohnQinAMD merged 1 commit into
mainfrom
fix/eagle-draft-greedy-topk

Conversation

@JohnQinAMD

@JohnQinAMD JohnQinAMD commented Sep 11, 2026 •

Copy link
Copy Markdown
Collaborator

Problem

sample_draft_proposal decides greedy-vs-sample from temperature alone. SamplingParams rewrites temperature=0 to temperature=1.0, top_k=1, so a greedy request is indistinguishable from a genuine T=1 one at that point. The draft then samples a sharp-but-not-degenerate distribution and proposes a non-argmax token often enough to break the chain, costing accept length.

Fix

Pass the per-request top_ks through and let a greedy row (top_k <= 1) propose its argmax.

This stays unbiased: eagle_sample renormalises the target by the same per-row top_ks before the accept test, so a greedy row's p is one-hot — X equal to the target argmax accepts, anything else rejects and the residual it resamples from is p itself. Both arms commit the target argmax.

The CUDA graph runner has to carry top_ks in a device buffer the way it already carries temperatures. Its synthetic SamplingBatchInfo used a host-side placeholder, so without the producer half the correction never sees a real top_k and only about a fifth of the loss comes back. The buffer is filled with TOP_K_ALL rather than -1: -1 is not a top_k this pipeline ever carries, and it would read as top_k <= 1, i.e. greedy, for the padded rows.

Usage

No flag. The fix is inside the EAGLE draft proposal, so it applies to any server already running EAGLE with rejection sampling, which is where the regression lives:

python3 -m sglang.launch_server \
  --model-path /path/to/GLM-5.2-MXFP4 --trust-remote-code \
  --tp 4 --ep-size 4 --kv-cache-dtype fp8_e4m3 \
  --dsa-prefill-backend tilelang --dsa-decode-backend tilelang \
  --max-running-requests 8 --cuda-graph-max-bs-decode 8 \
  --speculative-algorithm EAGLE --speculative-num-steps 5 \
  --speculative-eagle-topk 1 --speculative-num-draft-tokens 6

It only moves anything for greedy requests (temperature=0, which SamplingParams rewrites to temperature=1.0, top_k=1). Requests that genuinely sample are unaffected: their top_k is TOP_K_ALL, the greedy branch does not fire, and the proposal is the same draw as before.

Verify from accept len in the server's decode log under a greedy workload — it should sit near the no-speculation baseline rather than about 20% below it:

Decode batch, #running-req: 8, ..., accept len: 3.86, ...

Note that the CUDA graph half is what makes this visible. With --disable-cuda-graph the real top_ks already reached the proposal, so a server run that way shows the fixed accept length with or without this patch.

Results

GLM-5.2-MXFP4, MI355X, TP4/EP4, EAGLE steps=5 topk=1 draft=6, --speculative-use-rejection-sampling, measured against this repo's main.

GSM8K 200q 5-shot temperature=0, concurrency 8 — accept length is the metric this fix moves:

accept len accuracy invalid
main, rejection sampling off 3.778 0.945 0.000
main, rejection sampling on 3.080 0.930 0.000
+ this PR, rejection sampling on 3.882 0.940 0.000

Turning rejection sampling on costs 18% of accept length; this restores it to the level the argmax path already had.

bench_serving, random 1024/512, 12 prompts, concurrency 1. Read mean TPOT, not median: speculative decoding makes the per-token distribution bimodal, and the median sits inside the accepted-step mode where this fix does not show.

mean TPOT output tok/s
main, rejection sampling off 3.04 ms 310.2
main, rejection sampling on 3.29 ms 285.7
+ this PR, rejection sampling on 3.08 ms 303.7

Three repeats of that benchmark on one unchanged server give 310.2 / 312.9 / 309.7 tok/s and mean TPOT 3.02 / 3.00 / 3.03, so the noise floor is about 0.5% and these deltas sit well outside it.

The GSM8K accuracy column does not: three repeats on one unchanged server give 0.940 / 0.945 / 0.920, which spans every number in the table above. Nothing here is an accuracy claim in either direction — the point is that none of the arms produce invalid output.

With --speculative-use-rejection-sampling off, which is the default, sample_draft_proposal never runs and this patch changes nothing.

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>
@JohnQinAMD
JohnQinAMD force-pushed the fix/eagle-draft-greedy-topk branch from 360fe29 to 72ea484 Compare September 11, 2026 17:40
@JohnQinAMD
JohnQinAMD changed the base branch from dev_glm52_0907 to main September 11, 2026 17:40
@JohnQinAMD
JohnQinAMD merged commit 38cadae into main Sep 11, 2026
167 of 223 checks passed
Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Labels

None yet

Projects

None yet

Development

Successfully merging this pull request may close these issues.

1 participant