Skip to content

[Bugfix][DSv4] Make the C128A decode topk row stride capture-stable - #41

Merged
jasl merged 2 commits into
jasl:ds4-sm120-preview-devfrom
tobymao:fix/dsv4-c128a-capture-stable-stride
Aug 6, 2026
Merged

jasl merged 2 commits into
jasl:ds4-sm120-preview-devfrom
tobymao:fix/dsv4-c128a-capture-stable-stride

Conversation

@tobymao

@tobymao tobymao commented Aug 6, 2026

Copy link
Copy Markdown

[Bugfix][DSv4] Make the C128A decode topk row stride capture-stable

Fixes the intermittent output corruption on DeepSeek-V4-Flash under concurrent
long-context load (vllm-project#41834, our 4x GB10 / TP=4 / 1M-context
reports): mid-generation <|begin_of_sentence|> bursts, multilingual token
salad, terminal repetition — one request per batch, never the first.

Root cause

_build_c128a_metadata derives the row stride of the flat persistent
c128a_topk_buffer from the batch's max_seq_len:

active_topk_width = min(max(next_pow2(cm.max_seq_len // 128), 128), c128a_max)

The decode consumers of that layout (_forward_decode in the FlashMLA /
FlashInfer SM120 backends) run inside FULL cudagraphs. Cudagraph capture
builds attention metadata with max_seq_len = max_model_len
(gpu_model_runner.py, for_cudagraph_capture branch), so the captured
kernels bake the widest stride — 8192 for a 1M-context model — while the
runtime builder re-lays rows out at whatever the current batch gives
(128…4096 in practice).

Row 0 lines up at offset 0 under any stride and is always read correctly.
Every later decode row is read from stale bytes. The sharpest case: two decode
rows at runtime stride 4096 put decode_end at 8192, so a mixed batch
writes prefill row 0 at exactly the offset the captured kernel reads as decode
row 1
— local compressed indices [0..n-1, -1, …] that the decode kernel
treats as global slot ids. The victim then attends over the first pages of the
compressed KV pool: other requests' KV (→ token salad), or uninitialized fp8
(~2/256 random bytes decode to NaN → an all-NaN logits row; argmax over
all-NaN returns index 0, and token 0 is BOS — the BOS-burst symptom). Once NaN
enters the victim's hidden state it is written into its KV and every following
step is poisoned, which is why the corruption persists to end-of-stream.

This one mechanism explains the full observation matrix:

observation explanation
victim is never batch row 0 (150/150 NaN hits on row 1) row 0 reads correctly under any stride
single-request traffic never corrupts (~40 quiet trials) no row 1 exists
--enforce-eager clean (0 NaN events / 20 streams) no baked stride; builder and consumer agree every step
onset mid-generation, after the partner stream's prefill chunks each mixed batch rewrites the stale region row 1 reads
deterministic across replays, passed 13 determinism probes same baked addresses, same stale bytes
clean on other labs' configs widths match when max_model_len is close to the tested context
92k contexts much milder than 348k reads seq/128 stale entries: 719 vs ~2700

Evidence (measured, not inferred)

  • Builder-side width recording on the live engine: capture builds at width
    8192 (exactly 5 of them = the 5 FULL decode capture shapes), runtime
    builds at 128…4096 — every decode replay ran with a builder/consumer
    stride mismatch (hist={128: 712, 8192: 5, 256: 4, …, 4096: 12592}).
  • An in-graph trace buffer — sync-free captured device ops that the FULL
    graphs re-execute at replay — recorded per-layer nonfinite (count, row
    bitmask) pairs on the very forward whose output row was NaN. At the corrupt
    moment (2 decode rows, 2×~360k-token streams): embedding and layers 0–6
    fully finite, then L7.h = 4096@[row 1] — the entire hidden row goes
    NaN at the output of layer 7, a compress-ratio-128 layer, and stays NaN
    through the final norm. 38 of 39 recorded events were born at L7, one at
    L11 — only ratio-128 layers, only row 1; row 0 was healthy throughout.
  • The buffer readout at the same step, row 1 (lens = [2823, 2823, …]):
    • what the builder wrote at stride 4096: [140, 141, 138, 139, …],
      2823 valid slot ids, 0 negative — the victim's true top-k;
    • what the captured kernel reads at stride 8192:
      [61946, 61947, 61944, 61945, …], a stale row with 2209 foreign slot
      ids and 614 interleaved -1s — pages the victim does not own.

Fix

Lay decode rows out at c128a_max_compressed always. The only cost is wider
-1 padding writes in the build kernel (~1–2 ms per 16k-token prefill chunk);
kernel reads stay bounded by decode_lens via topk_length, so decode-side
work is unchanged.

Verification

  • Regression gate tests/v1/attention/test_dsv4_c128a_capture_stable_stride.py
    inspects the assignment (same spirit as the short-extend tiering gate):
    verified to fail on the unfixed tree and pass on the fixed one.
  • Fix-arm positive control on the live cluster (same boot, fix applied as a
    one-line mount): wcap=8192 wrun=8192 MISMATCH=False, width histogram
    single-valued at 8192, and the two candidate row-1 layouts read
    byte-identically — the captured kernels now read exactly what the builder
    wrote.
  • Battery A/B on 4x GB10 / TP=4 / 1M context / spec-off / FULL_AND_PIECEWISE,
    2x300k-token concurrent streams per round, engine-side NaN detection at the
    sampling-row gather (fires even when corruption never reaches scoreable
    text). Unfixed: 39 NaN events in round 1 (historical rate
    ~25%/stream/round). Fixed: 8 rounds / 16 streams, ZERO engine-side NaN
    events across ~105,000 decode steps
    (GRAPH-TRACE:HIT = 0,
    GATHER-CHK NaN = 0, NAN-ORIGIN = 0), zero leaked specials, zero
    multi-script output, and the builder width pinned at 8192 for all 105,014
    metadata builds of the run. At the pre-fix per-stream rate, 16 clean
    streams is p ≈ 0.01 by text scoring alone; the engine-side counter fired
    20–400 times per corrupt round pre-fix, so the discriminating power is far
    higher.

Relationship to d8885a3 (rejection-sampler recovered-token fix)

Complementary, not overlapping. That fix stops the rejection sampler from
emitting token 0 when handed an all-NaN draft row — downstream NaN
handling on the spec-decode path. This PR fixes the mechanism that computes
NaN rows inside the target model's forward with spec decode OFF (no
drafter, no recovered sampling — the configuration jasl noted "is a
different mechanism by construction"). It plausibly also explains where the
all-NaN draft rows come from in spec-on configs: the DSpark drafter runs the
same sparse attention path, and under spec decode every request contributes
multiple decode rows per step, i.e. more rows beyond row 0 reading the
buffer at the wrong stride.

AI assistance was used throughout (instrumentation, analysis, and the patch);
the change and its verification were reviewed end-to-end by the submitter.
This does not duplicate any open PR (jasl/vllm has none open; d8885a3
touches sample_recovered_tokens_kernel, a different subsystem).

The C128A builder derives active_topk_width from the batch's max_seq_len
(next_power_of_2(max_seq_len / compress_ratio)) and lays decode rows out in
the flat persistent c128a_topk_buffer at that stride. The decode consumers
(FlashMLA / FlashInfer SM120 _forward_decode) run inside FULL cudagraphs, and
cudagraph capture builds metadata with max_seq_len = max_model_len, so the
captured kernels bake the widest stride (8192 for a 1M-context model) while
the runtime builder writes rows at whatever the current batch's stride is.

Row 0 lines up at offset 0 under any stride and is always read correctly.
Every decode row after it is read from stale bytes: with two decode rows at a
runtime stride of 4096, the captured kernel reads row 1 at flat[8192:], which
is exactly where a mixed batch writes prefill row 0 -- local compressed
indices [0..n-1, -1, ...] that the decode kernel then treats as global slot
ids, attending over the first pages of the compressed KV pool (other
requests' KV, or uninitialized fp8 where random bytes decode to NaN).

Observed in production on 4x GB10 / TP=4 / 1M max_model_len as exactly this
signature: one request per concurrent batch -- never the first -- emitting
all-NaN logits rows (argmax over an all-NaN row returns token 0, i.e. BOS
bursts) or multilingual token salad, onset mid-generation after the partner
stream's prefill chunks, --enforce-eager always clean, single-request traffic
never affected (no row 1). An in-graph trace buffer written by kernels
captured inside the FULL graphs confirmed the builder/consumer stride
mismatch on the live engine (capture builds at width 8192, runtime builds at
128..4096).

Fix: the decode row stride must not depend on the batch; always lay rows out
at c128a_max_compressed. Cost is only wider -1 padding writes in the build
kernel (~1-2 ms per 16k-token prefill chunk); reads stay bounded by
decode_lens via topk_length.

The gate inspects the assignment rather than asserting on a computed value,
in the spirit of the short-extend tiering gate: verified to FAIL on the
unfixed tree and pass on the fixed one.

Co-authored-by: Claude <noreply@anthropic.com>
…resses must not move

Builds metadata once at max_seq_len = max_model_len (what cudagraph capture
does) and once for a small batch, then asserts decode row 1's stride and
address did not change. On the unfixed tree it fails with 'stride changed
with the batch's max_seq_len (8192 at capture vs 128 at runtime)'; verified
both ways on GB10 against the b2bdfc7 pin.

Co-authored-by: Claude <noreply@anthropic.com>
@tobymao

tobymao commented Aug 6, 2026

Copy link
Copy Markdown
Author

Pushed one addition: a GPU reproduction test (test_c128a_decode_row_addresses_survive_batch_width_changes) — builds metadata once at max_seq_len = max_model_len (what capture does) and once for a small batch, and asserts decode row 1's stride and address did not move. Fails on the unfixed pin with stride changed with the batch's max_seq_len (8192 at capture vs 128 at runtime); passes with the fix. Verified both ways on GB10.

Also filed the same fix upstream (the identical derivation is on vllm-project/vllm main): vllm-project#51318.

@jasl
jasl merged commit c21ad6b into jasl:ds4-sm120-preview-dev Aug 6, 2026
@jasl

jasl commented Aug 6, 2026

Copy link
Copy Markdown
Owner

Merged — thank you for a fix whose receipts made review easy. Validation before landing: your AST + GPU stride tests plus the restored tiering suite (6 green), the targeted sparse/attention suites (57 green; the 2 masked-MHA failures reproduce identically on the pre-fix tree — environmental), and two full serve cycles plus the standard gates and a benchy row against yesterday's baseline (pp identical, tg within the documented spread). Branch convergence was a welcome side effect: your branch carried #40's content, so landing it brought codex/ds4-sm120-min-enable and ds4-sm120-preview-dev back to SHA-equal.

One data point on your closing suspicion (that this also explains our drafter's all-NaN draft_probs rows): it doesn't — but the real culprit is your bug's capture-independent cousin. On a tree WITH your fix, --enforce-eager (no baked strides by construction) still produced ~159 NaN draft rows per 20-request run. Chased to ground: token_to_req_indices() caches its token→request mapping on the CommonAttentionMetadata object, and the drafter's multi-step loop rewrites that same object in place from the ragged first-pass layout to per-request rows — a stale hit hands draft rows > 0 an earlier request's identity in mixed batches, the SWA window anchors one slot past that request's last written token, and the drafter attends over unwritten fp8. Same disease family — producer and consumer disagreeing about row layout — different vector. Fixed in 2649427; counters went 159 → 0 in both eager and FULL-cudagraph runs. Your fix and that one now cover the layout-desync family from both ends.

tobymao added a commit to tobymao/vllm that referenced this pull request Aug 6, 2026
The C128A builder derives active_topk_width from the batch's max_seq_len
(next_power_of_2(max_seq_len / compress_ratio)) and lays decode rows out in
the persistent decode buffer at that stride. The decode consumers (FlashMLA /
FlashInfer SM120 _forward_decode) run inside FULL cudagraphs, and cudagraph
capture builds metadata with max_seq_len = max_model_len, so the captured
kernels bake the widest stride (8192 for a 1M-context model) while the
runtime builder writes rows at whatever the current batch's stride is.

Row 0 lines up at offset 0 under any stride and is always read correctly.
Every decode row after it is read at the wrong offset: stale slot ids from
earlier, differently-strided builds, which the decode kernel then treats as
valid global slots -- attending over compressed-KV pages the request does
not own (multilingual token salad) or over uninitialized fp8 where random
bytes decode to NaN. One NaN turns the whole logits row NaN, and argmax over
an all-NaN row returns index 0: the begin-of-sentence bursts reported in
vllm-project#41834.

Root-caused on 4x GB10 / TP=4 / 1M max_model_len with an in-graph trace
(sync-free device ops captured into the FULL graphs, so replays record what
they compute): NaN born only at compress-ratio-128 layers, only in rows >= 1,
with the builder's row 1 holding valid slot ids while the captured kernel
read a stale row. A/B under the reproducing load (2x300k-token concurrent
streams): 39 engine-side NaN events in round 1 unfixed; zero events across
8 rounds / 16 streams / ~105k decode steps fixed. Same fix merged in the
GB10 deployment fork as jasl#41.

Fix: the decode row stride must not depend on the batch; always lay rows out
at c128a_max_compressed. Cost is only wider -1 padding writes in the build
kernel; reads stay bounded by decode_lens via topk_length.

The gate inspects the assignment rather than asserting on a computed value:
verified to FAIL on the unfixed tree and pass on the fixed one.

Signed-off-by: tobymao <toby.mao@gmail.com>
Co-authored-by: Claude <noreply@anthropic.com>
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.

2 participants