[BUGFIX][Mamba][Qwen3.5] Zero freed SSM cache blocks on GPU - #35219
Conversation
Signed-off-by: Vadim Gimpelson <vadim.gimpelson@gmail.com>
There was a problem hiding this comment.
Code Review
This pull request introduces a bug fix for Mamba-based models, specifically addressing an issue where freed SSM cache blocks on the GPU were not being zeroed out. This could lead to incorrect state being reused in subsequent computations. The fix involves implementing a mechanism to track SSM blocks that are truly freed (i.e., their reference count drops to zero) during each scheduling step. These freed block IDs are then passed to the worker, which explicitly zeroes out the corresponding state tensors on the GPU. The changes are well-contained and correctly implemented across the scheduler and worker components, ensuring that Mamba's stateful cache is properly managed. The logic for identifying and collecting freed blocks is soundly integrated into the existing KV cache management lifecycle methods.
|
Could you please provide data on possible perf overhead? also with async scheduling I think it may be risky to zero on free, we may need to move this into the model runner to ensure it ends up in the correct order in the GPU stream |
Actual zeroing happens in |
|
Below is from slack discussion. But I think it's worth attention and sharing here Seems both FlashAttn and trt-llm attn use mul by 0 to mask not used values. |
|
How does this interact with prefix caching? If we zero out blocks when their ref_cnt hits zero, doesn't that mean they can't be re-used if something comes along later that gets a cache hit? Wouldn't it to be better to detect the event when we use a block for attention that was previously used for mamba (in some other dtype) and zero it out at that point? |
Ran on B200 With changes and without changes the Output Tokens vary around 15000+-300. Definitely nothing dramatical from perf point of view but not exact numbers. |
Good catch I broke the prefix caching :/ |
Signed-off-by: Vadim Gimpelson <vadim.gimpelson@gmail.com>
70dafd6 to
152ccc2
Compare
I redo this PR similar to what @tdoublep proposed. I decided to zero out every new block, whether it comes from attention or from the SSM. Justification: attention can also produce NaNs in certain corner cases. Getting garbage for one specific request is likely acceptable, but without zeroing, the NaN could propagate to all requests. |
|
pls take a look |
|
|
||
| def _zero_block_ids(self, block_ids: list[int]) -> None: | ||
| """Zero the raw KV cache memory for the given block IDs.""" | ||
| for raw_tensor, page_size in self.kv_cache_raw_buffers: |
There was a problem hiding this comment.
Would it be more efficient to build an index tensor and have one op to zero at all the block id slots?
There was a problem hiding this comment.
How many block ids would we normally see for a typical prefill/decode? Is it very few?
There was a problem hiding this comment.
This zeroing takes small amount of time. We do it once per forward step and only for new.
@benchislett Can you say right away does it code works in sync or async part?
There was a problem hiding this comment.
I notice that this is not specific to SSM blocks, and it clears all new KV blocks. Will this have a detrimental effect on prefills for non-mamba deployments where block_size=16?
In this case if we get a prefill of 8k tokens, that will be 512 new blocks, right? I think that would lead to 512 kernel invocations in this implementation. If that is indeed the case, this will not suffice.
There was a problem hiding this comment.
right,
I am optimizing it
There was a problem hiding this comment.
does it make sense to use torch.tensor for block ids and use a gpu operation to zero the indices in tensors?
There was a problem hiding this comment.
See my comment below
There was a problem hiding this comment.
does it make sense to use torch.tensor for block ids and use a gpu operation to zero the indices in tensors?
I implemented zeroing as a triton kernel
There was a problem hiding this comment.
Pls lets me know if there is a better way to do it
…ject#35219) Signed-off-by: Vadim Gimpelson <vadim.gimpelson@gmail.com>
…ject#35219) Signed-off-by: Vadim Gimpelson <vadim.gimpelson@gmail.com>
…ject#35219) Signed-off-by: Vadim Gimpelson <vadim.gimpelson@gmail.com> (cherry picked from commit 03a1823)
…ject#35219) Signed-off-by: Vadim Gimpelson <vadim.gimpelson@gmail.com>
PR vllm-project#35219 records every newly allocated full-attention/MLA block id into SingleTypeKVCacheManager.new_block_ids, but the scheduler only drains it via take_new_block_ids() when needs_kv_cache_zeroing, which equals has_mamba_layers. Models without Mamba layers therefore never drain the list, so it grows without bound and leaks host memory under sustained load (one int per allocated block per request). gc.freeze() at EngineCore startup excludes the list from gc.get_objects()/tracemalloc, which makes the growth easy to miss. Drain the per-step block ids unconditionally in the scheduler and only use them when zeroing is enabled. This bounds the list for all models without adding a constructor flag or reading needs_kv_cache_zeroing twice; for Mamba models the drain already happened in that branch, so their behavior is unchanged. Fixes vllm-project#44175 Signed-off-by: Ting Sun <suntcrick@gmail.com>
…eadout buffer DEBUG / INVESTIGATION COMMIT -- searchable keywords: !!!! exclamation NaN logits argmax token-0 mamba zeroing block recycling GDN readout fp16 overflow sq_intr `!!!!!!` in the output is token id 0 (`!` in the Qwen vocabulary), returned by argmax over a logits row that is entirely NaN. Two INDEPENDENT causes produce that same row; fixing one does not fix the other. Upstream vllm-project#55291 lists both. -------------------------------------------------------------------------- 1) Recycled KV blocks handed to a Mamba group were never zeroed -------------------------------------------------------------------------- vllm/v1/core/single_type_kv_cache_manager.py A KV block is one shared tile: cache groups alias the same bytes, as KVCacheTensor states ("cache groups overlay each other, which is sound because a block ID is owned by one group at a time"). `needs_kv_cache_zeroing` exists to clean that handover, but only armed when the NEW owner is an AttentionSpec group. A block that attention frees and Mamba then allocates therefore begins life reading quantized int8 KV bytes as recurrent state -> NaN/Inf in the state -> NaN logits. TWO sites are required. The second is load-bearing: in `mamba_cache_mode=align` MambaManager.allocate_new_blocks() calls block_pool.get_new_blocks() directly and never reaches the base-class recording, so the isinstance gate alone is INERT. We shipped the one-line version by mistake and it ran 40 minutes doing nothing; only the test caught it. Does NOT contradict vllm-project#35219 for the reason it gives ("they overwrite their state fully each step"): that is false in `align` mode, which is later. The state lives in one block at a time and the precopy moves it, so blocks ahead of the write cursor -- including the num_speculative_blocks reserved for MTP -- are allocated and read before being written. Ordering is safe, verified in the deployed build: zeroing runs inside _update_states, i.e. before _prepare_inputs and before preprocess_mamba (the precopy), and before kv_cache_block_copies (the CoW fill). Block ids are pool-global (all groups share kv_cache_config.num_blocks), so a Mamba id can never index out of an attention tensor. Only blocks fresh from get_new_blocks() are recorded, so the speculative blocks that the `align` branch recycles within a request are untouched -- live state is never zeroed. -------------------------------------------------------------------------- 2) GDN readout overflows fp16 before RMSNormGated can normalize it -------------------------------------------------------------------------- vllm/third_party/flash_linear_attention/ops/fused_sigmoid_gating.py vllm/model_executor/layers/mamba/gdn/qwen_gdn_linear_attn.py The GDN recurrent state is an unnormalized accumulator, normalized only at the readout. With an fp32 SSM state the readout can exceed fp16's 65504, overflow to inf on the store, and RMSNormGated turns inf into NaN. Requires fp16 activations AND an fp32 state. bf16 shares fp32's exponent range and is immune, which is why this does not bite everyone. Forcing fp32 on the mamba cache does NOT help -- that is the precondition, not the cure. Upstream PR vllm-project#54146 alone is NOT sufficient in this tree: it widens `o` inside the kernel wrapper, but the value is copied straight into `core_attn_out`, which is allocated in the activation dtype, and the norm runs after that copy. The overflow simply moves one line down. That is the concrete reason vllm-project#55291 reports vllm-project#54146 as "not confirmed as sufficient". Hence the companion change here: _core_attn_out_dtype() + the five buffer allocation sites. It returns fp32 only for (fp16 activations, fp32 state); bf16 and fp32 stay byte-identical. RMSNormGated already does x.float() internally, so it consumes fp32 unchanged. `z` is deliberately NOT widened: the kernels write it in the activation dtype. -------------------------------------------------------------------------- ATTRIBUTION / PROVENANCE -- read before opening any PR -------------------------------------------------------------------------- * The fused_sigmoid_gating.py hunk and the test_fused_sigmoid_gating_readout_finite_under_large_fp32_state test are PORTED VERBATIM from upstream vllm-project/vllm PR vllm-project#54146. NOT authored here. Same fix also filed at ai-infos#22. Any PR carrying this must depend on and cite vllm-project#54146, never present it as original work. * Everything else in this commit is ours: the Mamba zeroing change, the core_attn_out widening, and their tests. -------------------------------------------------------------------------- VALIDATION -------------------------------------------------------------------------- Tests (CPU only, no GPU): tests/v1/core/test_single_type_kv_cache_manager.py 3 new, pass; the recycling one FAILS without the two-site change tests/kernels/mamba/test_gdn_readout_buffer_dtype.py 7 pass; two of them demonstrate the defect itself (1e6 stored into an fp16 buffer -> inf -> NaN after the norm; the same value in a wide buffer survives) Cause 1 measured in production on 4x RX 7900 XTX (gfx1100), Qwen3.5-MoE hybrid, TP4, kv_cache_dtype=int8_per_token_head, prefix caching, chunked prefill, MTP. Exposure = prompt tokens not served from cache + generated tokens, because the defect is driven by block recycling, not by time. Two independent instruments (the Prometheus counters and the integral of the logged prompt throughput) agree to 0.012%. run 3 no patch 1,075,538 exposure -> NaN, engine dead run 5 no patch 151,867 exposure -> NaN, engine dead run 7 PATCHED 4,245,144 exposure -> clean, engine alive ~1 in 1000 under the null that the patch does nothing. Baseline is n=2: quote it as "1 in 1000 with a two-sample baseline", never as a decimal. The control arm (unpatched node) then died of the same NaN on 2026-09-08 07:23:35 with 248,320 NaN entries, at running=1 and 3.4% cache usage -- so this needs block recycling, NOT high concurrency. Cause 2 is NOT validated in production. It is written, tested, and deployed nowhere. -------------------------------------------------------------------------- STILL OPEN (do not confuse with the above) -------------------------------------------------------------------------- Workers stop responding with corrupted_requests_total == 0 and an amdgpu `sq_intr` minutes earlier. Different failure: no NaN, nothing to dump. 57 events in 7 bursts over 77 days, confined to 2 of one node's 4 GPUs and zero on the other node, with dates predating all of this work. No fix, no cause. If corrupted_requests_total does NOT rise, it is that one, not this. Signed-off-by: JartX <sagformas@epdcenter.es>
KDA mamba recurrent/conv state blocks are never zeroed when recycled, so a new request inherits a finished request's finite-but-wrong state -> con=32 NIAH recall drops to ~65%. needs_kv_cache_zeroing is True for mamba (issue vllm-project#35219) but two isinstance(AttentionSpec) gates exclude MambaSpec. Add a separate pure-torch zero-on-recycle channel for mamba blocks (V2 runner update_requests).
Two independent bugs made a restored hybrid prefix produce garbage; both can only happen on a hybrid (Mamba + attention) model, which is why the pure-attention model was clean all along. Block indexing: the scheduler addresses KV blocks -- num_blocks logical pages of block_size (528) tokens -- but attention kernels lay each page out as 'kv_cache_spec.block_size // kernel_block_size' (528 // 16 = 33) consecutive kernel blocks, so tensor dimension 0 is a finer ID space. The connector indexed it with scheduler block IDs, so a save or restore touched 1/33 of a page at the wrong offset. register_kv_caches now re-views every page as one row, putting dimension 0 back in the scheduler's ID space and keeping the kernel layout in the trailing dimensions; the view is free and nothing is copied. Load ordering: vLLM zeroes freshly allocated attention blocks on the compute stream (KVBlockZeroer, hybrid only, because attention and fp32 SSM states share one block pool -- see vllm-project/vllm#35219). KVFlow copies on a private stream with no dependency on it, so an H2D load could be issued before that zeroing and be erased by it. get() now waits for the compute stream, a device-side dependency; the blocking CPU sync stays opt-in behind IAXL_CACHE_STREAM_SYNC_ON_GET. Also drop two redundant prefill gates in build_connector_meta and return early from wait_for_save when nothing is queued. Verified on 2x RTX 4090, TP2: 102 unit tests pass and all 9 GPU gates pass. GSM8K (8-shot, thinking off), 200 questions: cold 0.935 -> warm 0.920, against 0.835 warm with only the block-indexing half fixed and 0.710 with neither. TTFT 670 ms -> 148 ms.
Two independent bugs made a restored hybrid prefix produce garbage; both can only happen on a hybrid (Mamba + attention) model, which is why the pure-attention model was clean all along. Block indexing: the scheduler addresses KV blocks -- num_blocks logical pages of block_size (528) tokens -- but attention kernels lay each page out as 'kv_cache_spec.block_size // kernel_block_size' (528 // 16 = 33) consecutive kernel blocks, so tensor dimension 0 is a finer ID space. The connector indexed it with scheduler block IDs, so a save or restore touched 1/33 of a page at the wrong offset. register_kv_caches now re-views every page as one row, putting dimension 0 back in the scheduler's ID space and keeping the kernel layout in the trailing dimensions; the view is free and nothing is copied. Load ordering: vLLM zeroes freshly allocated attention blocks on the compute stream (KVBlockZeroer, hybrid only, because attention and fp32 SSM states share one block pool -- see vllm-project/vllm#35219). KVFlow copies on a private stream with no dependency on it, so an H2D load could be issued before that zeroing and be erased by it. get() now waits for the compute stream, a device-side dependency; the blocking CPU sync stays opt-in behind IAXL_CACHE_STREAM_SYNC_ON_GET. Also drop two redundant prefill gates in build_connector_meta and return early from wait_for_save when nothing is queued. Verified on 2x RTX 4090, TP2: 102 unit tests pass and all 9 GPU gates pass. GSM8K (8-shot, thinking off), 200 questions: cold 0.935 -> warm 0.920, against 0.835 warm with only the block-indexing half fixed and 0.710 with neither. TTFT 670 ms -> 148 ms.
Essential problem
Fixes #35138
Workaround for Dao-AILab/flash-attention#1974
Hybrid models (e.g. Qwen3.5-397B-A17B) share a unified block pool between attention (fp8/fp16) and Mamba/SSM (fp32) layers. When a block previously used by Mamba (fp32 state) is reallocated to an attention layer with a smaller dtype, leftover fp32 bit patterns can appear as NaN/Inf in the new dtype. Attention kernels (FlashAttn3, FlashInfer-TRTLLM, etc.) use multiply-by-zero masking for unused positions, which does not clear NaN (
0 * NaN = NaN). The stale NaN then propagates across all requests sharing the same KV-cache block, causing progressive accuracy degradation over time.What this PR does
Zeroes GPU memory of freshly allocated full-attention KV-cache blocks before they are used, but only for hybrid models (models with Mamba layers). Mamba/SSM blocks are not zeroed (they overwrite their state fully on each step). The approach:
Scheduler side —
SingleTypeKVCacheManagertracks block IDs allocated since the last scheduling step (only forFullAttentionSpeclayers). After scheduling, the scheduler drains these IDs intoSchedulerOutput.new_block_ids_to_zero, gated behindself.has_mamba_layers.Worker side —
GPUModelRunner._update_states()receives the block IDs and calls_zero_block_ids(), which launches a single Triton kernel (_zero_kv_blocks_kernel) to zero the corresponding memory across all KV-cache segments in one GPU launch.Optimized zeroing — A one-time
_init_kv_zero_meta()precomputes absolute byte addresses of all KV-cache segments (handling both block-dim-0 and block-dim-1 layouts, multi-buffer backends, and virtual block splitting). Block IDs are transferred via pre-allocated pinned memory to overlap the H2D copy with kernel launch. This avoids 15 separateindex_fill_calls (For Qwen3.5-379B, one per layer).CuMem compatibility —
_init_kv_zero_meta()is called ingpu_worker.pyoutside the CuMem pool context, so the bookkeeping tensors (segment addresses, block-ID buffers) use the standard PyTorch allocator and survive sleep/wake cycles.Scope
self.has_mamba_layersin the scheduler. Non-hybrid (pure attention) models are completely unaffected.FullAttentionSpecblocks — Mamba/SSM blocks are not zeroed (they overwrite their state fully each step); only attention blocks that may inherit stale Mamba fp32 data are cleared.Performance overhead
Measured on B200, Qwen/Qwen3-0.6B, BS=500:
End-to-end benchmark on B200 (Qwen3.5-397B-A17B-FP8, TP=1 PP=1 DP=8, 2048 prompts, 500 output tokens) showed no measurable throughput degradation (output tokens/s within ±2% noise).
Test plan