Skip to content

[PD][Spec] Give a PD decode's first EAGLE draft step a proposal distribution - #41783

Open
salexspb wants to merge 4 commits into
sgl-project:mainfrom
salexspb:upstream/pd-eagle-draft-proposal-20260929
Open

salexspb wants to merge 4 commits into
sgl-project:mainfrom
salexspb:upstream/pd-eagle-draft-proposal-20260929

Conversation

@salexspb

@salexspb salexspb commented Sep 30, 2026 •

Copy link
Copy Markdown
Contributor

PD stack

  1. [PD] Wait for prefill completion before Mooncake early KV transfer — #41395.
  2. [PD][Spec] Give a PD decode's first EAGLE draft step a proposal distribution — this PR.

Problem

With PD and EAGLE, whenever rejection sampling is enabled, the verifier is handed each draft step's proposal distribution. ROCm enables rejection sampling by default for EAGLE top-k 1 (#37134, _should_auto_enable_hip_rejection_sampling), which is why this crashes there out of the box. EagleWorkerV2.draft_forward starts that list with spec_info.draft_probs. PD prefill sends decode the first draft token, its probability and its hidden state, but not the distribution, so build_eagle_disagg_draft_input leaves draft_probs unset. The PD decode then fails on its first forward:

eagle_worker_v2.py, in draft_forward
    torch.stack(draft_probs_list, dim=1)
TypeError: expected Tensor as element 0 in argument 0, but got NoneType

Under a draft CUDA graph, replay skips the proposal copy, and the verify would read a stale buffer instead. This happens with PD + EAGLE (including native MTP) whenever rejection sampling is enabled. Aggregated serving is unaffected.

Fix

build_eagle_disagg_draft_input sets draft_probs to a one-hot FP32 proposal at each received draft token (pd_first_draft_proposal) whenever rejection sampling is on. Single-layer EAGLE gets one (batch, vocab) row. Multi-layer EAGLE gets the (batch, num_steps, vocab) chain, because its verify reads the received chain directly and otherwise fails with "draft_probs missing".

The verify stays exact. It accepts the token X with probability p(X) and otherwise resamples from the target with X removed, so the committed token is distributed as p, the verifier's target distribution. A greedy request's target is one-hot, so it accepts X exactly when X is the target argmax, with or without the true proposal, and its acceptance is unchanged. Sampled requests accept X with probability p(X) rather than Σ min(p, q), and only at their first decode step. Nothing is added to the PD wire.

#40953 instead transfers the full FP32 distribution through PD metadata, which needs a bootstrap wire change and costs vocab × 4 bytes per request. This PR is the minimal fix that restores correctness; the two could be combined.

Tests

test/registered/unit/speculative/test_pd_eagle_draft_proposal.py (CPU, 5 tests):

  • the built draft input carries a one-hot FP32 proposal at each request's drafted token;
  • the proposal's width matches the target vocab, as eagle_sample requires;
  • there is no proposal without rejection sampling;
  • multi-layer EAGLE gets a one-hot FP32 proposal at each received chain step;
  • enumerating a single-token chain verify over a drafting q, the committed-token distribution equals p.

On this head: 5 passed. python -m compileall and import sglang are clean.

End-to-end validation (on a HiSparse fork build)

Before this change, a PD 4:4 decode (native MTP, EAGLE 3 steps, top-k 1, rejection sampling on) died at its first request with the torch.stack TypeError above. After it, in every configuration tested (aggregated TP8, PD 4:4 resident decode, PD 4:4 HiSparse decode), draft CUDA-graph capture completed, every service stayed up with no traceback, and 8/8 greedy chat prompts (including two needle prompts) were answered correctly, with mean accept length in the 2.6–2.65 range across configurations, matching the aggregated baseline.

These runs used single-layer EAGLE (native MTP). Multi-layer EAGLE PD is covered by the unit test only.


CI States

Latest PR Test (Base): ✅ Run #37887765090
Latest PR Test (Extra): ❌ Run #37887764918
Latest PR Test (AMD ROCm 10): ❌ Run #37887765201

@salexspb
salexspb force-pushed the upstream/pd-eagle-draft-proposal-20260929 branch 4 times, most recently from 044595f to ecbb734 Compare October 2, 2026 22:27
@salexspb
salexspb marked this pull request as ready for review October 2, 2026 22:53
@salexspb
salexspb force-pushed the upstream/pd-eagle-draft-proposal-20260929 branch 2 times, most recently from 2bf5a18 to c7e556a Compare October 2, 2026 23:17
…ibution

With rejection sampling enabled, each draft step hands the verifier the
distribution its token was drawn from (ROCm enables rejection sampling by
default for EAGLE top-k 1). PD prefill sends decode the first draft token,
its probability and hidden state, but not that distribution, so
build_eagle_disagg_draft_input left draft_probs unset. The first decode
forward then failed in EagleWorkerV2.draft_forward (torch.stack over a None
proposal); under a draft CUDA graph the stale proposal buffer would have
been verified instead.

Decode now proposes each received token with probability one: one row
per request for single-layer EAGLE, and one per chain step for
multi-layer EAGLE, whose verify reads the received chain directly and
otherwise fails on the missing proposal. The verify stays exact: it
accepts the token with probability p(X) and otherwise resamples from the
target without it. A greedy request's target is one-hot, so its
acceptance does not depend on the proposal.
@ShangmingCai
ShangmingCai force-pushed the upstream/pd-eagle-draft-proposal-20260929 branch from c7e556a to bffaf71 Compare October 6, 2026 08:35

@ShangmingCai ShangmingCai 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.

LGTM

@ShangmingCai

Copy link
Copy Markdown
Collaborator

/tag-and-rerun-ci

@github-actions github-actions Bot added the run-ci CI: run the baseline test suite on this PR label Oct 6, 2026
@ShangmingCai

Copy link
Copy Markdown
Collaborator

/rerun-failed-ci

@ShangmingCai ShangmingCai added the highest-priority CI: all three control labels, plus never batch-cancelled or stale-closed label Oct 8, 2026

This branch has not been deployed

No deployments
Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Labels

highest-priority CI: all three control labels, plus never batch-cancelled or stale-closed run-ci CI: run the baseline test suite on this PR

Projects

None yet

Development

Successfully merging this pull request may close these issues.

2 participants