Repository navigation
Conversation
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>
|
Pushed one addition: a GPU reproduction test ( Also filed the same fix upstream (the identical derivation is on vllm-project/vllm main): vllm-project#51318. |
|
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 One data point on your closing suspicion (that this also explains our drafter's all-NaN |
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>
[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 tokensalad, terminal repetition — one request per batch, never the first.
Root cause
_build_c128a_metadataderives the row stride of the flat persistentc128a_topk_bufferfrom the batch'smax_seq_len:The decode consumers of that layout (
_forward_decodein 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_capturebranch), so the capturedkernels 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_endat 8192, so a mixed batchwrites 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 kerneltreats 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;
argmaxoverall-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:
--enforce-eagerclean (0 NaN events / 20 streams)max_model_lenis close to the tested contextseq/128stale entries: 719 vs ~2700Evidence (measured, not inferred)
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}).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 goesNaN 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.
lens = [2823, 2823, …]):[140, 141, 138, 139, …],2823 valid slot ids, 0 negative — the victim's true top-k;
[61946, 61947, 61944, 61945, …], a stale row with 2209 foreign slotids and 614 interleaved
-1s — pages the victim does not own.Fix
Lay decode rows out at
c128a_max_compressedalways. The only cost is wider-1padding writes in the build kernel (~1–2 ms per 16k-token prefill chunk);kernel reads stay bounded by
decode_lensviatopk_length, so decode-sidework is unchanged.
Verification
tests/v1/attention/test_dsv4_c128a_capture_stable_stride.pyinspects the assignment (same spirit as the short-extend tiering gate):
verified to fail on the unfixed tree and pass on the fixed one.
one-line mount):
wcap=8192 wrun=8192 MISMATCH=False, width histogramsingle-valued at 8192, and the two candidate row-1 layouts read
byte-identically — the captured kernels now read exactly what the builder
wrote.
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, zeromulti-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).