From b8edada26059256276ef9e69887c2d6a751336f9 Mon Sep 17 00:00:00 2001 From: tobymao Date: Thu, 6 Aug 2026 09:30:34 -0700 Subject: [PATCH 1/2] [Bugfix][DSv4] Make the C128A decode topk row stride capture-stable 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 --- .../test_dsv4_c128a_capture_stable_stride.py | 88 +++++++++++++++++++ vllm/models/deepseek_v4/sparse_mla.py | 14 +-- 2 files changed, 95 insertions(+), 7 deletions(-) create mode 100644 tests/v1/attention/test_dsv4_c128a_capture_stable_stride.py diff --git a/tests/v1/attention/test_dsv4_c128a_capture_stable_stride.py b/tests/v1/attention/test_dsv4_c128a_capture_stable_stride.py new file mode 100644 index 000000000000..77c5d34da725 --- /dev/null +++ b/tests/v1/attention/test_dsv4_c128a_capture_stable_stride.py @@ -0,0 +1,88 @@ +# SPDX-License-Identifier: Apache-2.0 +# SPDX-FileCopyrightText: Copyright contributors to the vLLM project +"""Regression gate: the C128A decode topk row stride must be capture-stable. + +The C128A builder lays decode rows out in one flat persistent buffer at +``active_topk_width`` ints per row, and the decode consumers (FlashMLA / +FlashInfer SM120 ``_forward_decode``) run inside FULL cudagraphs, which bake +the row stride they saw at capture time. Capture builds metadata with +``max_seq_len = max_model_len``. If the runtime build derives the stride from +the batch's ``max_seq_len`` instead, the builder writes rows at a narrower +stride than the captured kernels read: row 0 still lines up at offset 0, but +every later decode row is read from stale bytes -- in a mixed batch, prefill +row 0 is written at exactly the offset the captured kernels read as decode +row 1. Observed in production as one request per batch (never the first) +emitting NaN-logits/BOS bursts or multilingual token salad under concurrent +long-context load, while ``--enforce-eager`` stays clean. + +Asserted by inspecting the assignment, in the spirit of the tiering gate: a +gate on a value the builder no longer computes cannot fail. +""" + +import ast +import inspect + + +def _active_topk_width_assignments() -> list[ast.Assign]: + from vllm.models.deepseek_v4 import sparse_mla + + tree = ast.parse(inspect.getsource(sparse_mla)) + build_fn = next( + node + for node in ast.walk(tree) + if isinstance(node, ast.FunctionDef) and node.name == "_build_c128a_metadata" + ) + return [ + node + for node in ast.walk(build_fn) + if isinstance(node, ast.Assign) + and any( + isinstance(t, ast.Name) and t.id == "active_topk_width" + for t in node.targets + ) + ] + + +def test_c128a_decode_stride_is_batch_independent(): + assigns = _active_topk_width_assignments() + assert assigns, ( + "_build_c128a_metadata no longer assigns active_topk_width; " + "re-point this gate at wherever the C128A row stride is computed" + ) + for node in assigns: + for sub in ast.walk(node.value): + attr = sub.attr if isinstance(sub, ast.Attribute) else None + name = sub.id if isinstance(sub, ast.Name) else None + assert (attr or name) != "max_seq_len", ( + "C128A row stride is derived from the batch's max_seq_len: " + f"`{ast.unparse(node)}`. FULL-cudagraph decode kernels bake " + "the capture-time stride (max_seq_len = max_model_len), so a " + "batch-dependent stride desynchronizes the builder's layout " + "from the captured readers for every decode row after row 0." + ) + + +def test_c128a_build_kernel_iterates_the_same_stride(): + """The build kernel's per-row iteration bound must be the row stride. + + ``build_c128a_topk_metadata`` writes rows at ``max_compressed_tokens`` + ints per row; passing anything other than ``active_topk_width`` would + desynchronize producer and consumer inside a single step. + """ + from vllm.models.deepseek_v4 import sparse_mla + + tree = ast.parse(inspect.getsource(sparse_mla)) + build_fn = next( + node + for node in ast.walk(tree) + if isinstance(node, ast.FunctionDef) and node.name == "_build_c128a_metadata" + ) + kws = [ + kw + for node in ast.walk(build_fn) + if isinstance(node, ast.Call) + for kw in node.keywords + if kw.arg == "max_compressed_tokens" + ] + assert kws, "build_c128a_topk_metadata call lost its max_compressed_tokens kwarg" + assert ast.unparse(kws[0].value) == "active_topk_width" diff --git a/vllm/models/deepseek_v4/sparse_mla.py b/vllm/models/deepseek_v4/sparse_mla.py index 1f34affe83f3..9d94706051d5 100644 --- a/vllm/models/deepseek_v4/sparse_mla.py +++ b/vllm/models/deepseek_v4/sparse_mla.py @@ -271,13 +271,13 @@ def _build_c128a_metadata( assert cm.positions is not None, ( "positions is required for C128A metadata build" ) - active_topk_width = min( - max( - triton.next_power_of_2(max(cm.max_seq_len // self.compress_ratio, 1)), - _C128A_TOPK_ALIGNMENT, - ), - self.c128a_max_compressed, - ) + # FULL-cudagraph decode kernels bake this layout's row stride at + # capture time (capture builds with max_seq_len = max_model_len), so + # the stride must not depend on the batch. A narrower runtime layout + # makes the captured kernels read decode rows r >= 1 from stale bytes: + # in a mixed batch, prefill row 0 is written at exactly the offset the + # captured kernel reads as decode row 1. + active_topk_width = self.c128a_max_compressed block_size = self.kv_cache_spec.block_size // self.compress_ratio global_decode, decode_lens, prefill_local = build_c128a_topk_metadata( cm.positions[:num_total], From e0ee8adf3377cf78a8c817969921aa08d0a744b3 Mon Sep 17 00:00:00 2001 From: tobymao Date: Thu, 6 Aug 2026 15:17:13 -0700 Subject: [PATCH 2/2] test(dsv4): GPU reproduction of the stride mismatch -- decode row addresses 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 b2bdfc7d31c9 pin. Co-authored-by: Claude --- .../test_dsv4_c128a_capture_stable_stride.py | 74 +++++++++++++++++++ 1 file changed, 74 insertions(+) diff --git a/tests/v1/attention/test_dsv4_c128a_capture_stable_stride.py b/tests/v1/attention/test_dsv4_c128a_capture_stable_stride.py index 77c5d34da725..81aa2b0d75da 100644 --- a/tests/v1/attention/test_dsv4_c128a_capture_stable_stride.py +++ b/tests/v1/attention/test_dsv4_c128a_capture_stable_stride.py @@ -86,3 +86,77 @@ def test_c128a_build_kernel_iterates_the_same_stride(): ] assert kws, "build_c128a_topk_metadata call lost its max_compressed_tokens kwarg" assert ast.unparse(kws[0].value) == "active_topk_width" + + +def test_c128a_decode_row_addresses_survive_batch_width_changes(): + """GPU reproduction of the corruption: decode row addresses must not move. + + FULL-cudagraph capture builds metadata with ``max_seq_len = + max_model_len``, and the captured decode kernels bake the resulting row + addresses. Every later build must therefore lay decode rows out at those + same addresses, or the captured kernels read rows r >= 1 from bytes the + builder no longer writes. On the unfixed tree this fails concretely: a + build at batch ``max_seq_len=300`` puts row 1 at stride 128 while the + capture-time build put it at stride 8192, so row 1's address moves by + (8192 - 128) * 4 bytes and the replayed kernel reads stale memory. + """ + import pytest + import torch + + if not torch.cuda.is_available(): + pytest.skip("requires CUDA") + + from types import SimpleNamespace + + from vllm.models.deepseek_v4.sparse_mla import DeepseekV4FlashMLAMetadataBuilder + + device = torch.device("cuda") + max_tokens = 64 + width_cap = 8192 + + b = DeepseekV4FlashMLAMetadataBuilder.__new__(DeepseekV4FlashMLAMetadataBuilder) + b.reorder_batch_threshold = 1 + b.compress_ratio = 128 + b.c128a_max_compressed = width_cap + b.kv_cache_spec = SimpleNamespace(block_size=256) + b.c128a_topk_buffer = torch.full( + (max_tokens, width_cap), -1, dtype=torch.int32, device=device + ) + b.c128a_decode_lens_buffer = torch.zeros( + max_tokens, dtype=torch.int32, device=device + ) + + def build(max_seq_len): + # Two single-token decode requests -- the smallest batch that has a + # row 1. The split helper's pure-decode path only needs the CPU + # bookkeeping fields plus is_prefilling for the tiering gate. + cm = SimpleNamespace( + query_start_loc_cpu=torch.tensor([0, 1, 2], dtype=torch.int32), + num_reqs=2, + max_query_len=1, + num_actual_tokens=2, + max_seq_len=max_seq_len, + is_prefilling=torch.tensor([False, False]), + positions=torch.tensor([256, 299], dtype=torch.int64, device=device), + block_table_tensor=torch.zeros((2, 8), dtype=torch.int32, device=device), + slot_mapping=torch.zeros(2, dtype=torch.int64, device=device), + ) + req_id = torch.tensor([0, 1], dtype=torch.int32, device=device) + fields = b._build_c128a_metadata(cm, req_id) + return fields["c128a_global_decode_topk_indices"] + + # Capture-time build (max_seq_len = max_model_len), then a runtime build + # for a small batch. + wide = build(1_048_576) + narrow = build(300) + + assert narrow.shape[-1] == wide.shape[-1], ( + f"decode topk row stride changed with the batch's max_seq_len " + f"({wide.shape[-1]} at capture vs {narrow.shape[-1]} at runtime): " + "FULL-cudagraph consumers baked the capture-time stride and will " + "read every decode row after row 0 from stale bytes" + ) + assert narrow[1].data_ptr() == wide[1].data_ptr(), ( + "decode row 1 moved between builds; captured kernels read it at the " + "capture-time address" + )