Repository navigation
[Kernel] Fix int32 page-offset overflow in the paged KV store kernels - #40846
tianxiaojiang4 wants to merge 1 commit into
Conversation
Fix verifiedSame escalating workload that reproducibly faults on an unpatched build, run against this change on 4x MI355X (gfx950, ROCm 7.2.4). An instrumented build prints the offset each expression computes, so the run shows the offset going past int32 rather than merely not crashing.
The patched run crossed the int32 boundary at batch 136 and served 7 more batches past it, ending 1.5x beyond the crossing with no fault and no change in pool occupancy (steady at 1%). Checklist update:
🤖 Generated with Claude Code |
Regression test verified against both treesThe test is only worth having if it fails on the unpatched code, so I ran it both ways on 4x MI355X (gfx950, ROCm 7.2.4), same container, only the mounted tree differing. Both kernels fault before the fix, so the indexer case is a genuine regression test and not just a correctness check at high indices — I had expected it might pass, given that the flashmla fp8 path survives its own overflow in a live server. Runtime is ~10 s per arm. The cost is memory, not time: 4.00 GiB and 8.00 GiB allocations, skipped when unavailable. 🤖 Generated with Claude Code |
2604c2f to
753872f
Compare
|
@kkHuang-amd @hnyls2002 @amd-danli103 — tagging you from This fixes an int32 overflow of the page offset in the two ROCm Triton KV store kernels. Why each of you:
Measured on MI355X with One ask: could someone with permission add the |
sgl-project#41159 fixed the int32 offset overflow in the Triton DSv4 KV store kernels but added no test, so nothing stops it returning. Covers both kernels (flashmla and indexer) at a slot whose page offset exceeds int32, and asserts the high-slot store does not clobber other pages. The cases skip unless the device has room for the 4 GiB / 8 GiB fixture. Co-Authored-By: Claude Opus 5 <noreply@anthropic.com>
753872f to
200e6b7
Compare
|
Closed in favor of #41159 |
Motivation
_triton_fused_store_flashmla_kerneland_triton_fused_store_indexer_kernelnarrow the KV slot index to int32 and then multiply it by the page stride:Triton wraps that multiply silently. Once the paged cache is addressed past 2^31 elements the store lands at a wrapped address and the GPU faults.
How it shows up
Measured on DeepSeek-V4.1-Flash, 4x MI355X (gfx950), TP4/EP4, page_size 256. The cache is
(196988, 149760)= 29.5 GiB, so the rope stride isbytes_per_page // 2= 74880 and int32 is exhausted at page 28679, i.e. slot index 7,341,824:page * 74880AMD_SERIALIZE_KERNEL=3attributes the fault to_triton_fused_store_flashmla_kernelon all four ranks. Without serialization the error surfaces at the next device sync, which makes it look like an unrelated failure elsewhere in the forward pass.Reproduced on ROCm 7.14 and 7.2.4, on two aiter builds, and on both a vendor image and
lmsysorg/sglang:dev-dsv41-mi35x.Why the radix cache is involved
Reaching the threshold needs prefix caching enabled: retained slots make the allocator issue monotonically climbing indices. With
--disable-radix-cacheslots are recycled and the index never approaches the crossing, which is why disabling the cache "fixes" it. The cache is not at fault; it just drives the index far enough to wrap.The trigger is ~7.34M cumulative distinct prefilled tokens and is independent of:
That invariance is what points at an index rather than a resource.
Impact
Once prefix caching retains slots the allocator stops reusing low indices, so the index tracks cumulative uncached tokens regardless of how full the pool is — an instrumented run reached slot 10,987,776 at 1% pool occupancy. Any workload that prefills ~7.34M uncached tokens crosses it.
For scale, on a long-context agentic replay (~150k-token prompts, ~0.95 hit rate) the rate of uncached tokens puts the crossing inside a single one-hour run from concurrency 8 upward:
A KV pool smaller than the crossing (7,341,824 slots at this geometry) cannot reach it, so a larger pool makes the fault more likely, not less. Disabling the prefix cache also avoids it, at the cost of all prefix reuse.
Scope
Two gates confine this, both at the dispatch site in
kernels/ops/attention/dsv4/attn.py:ROCm only, and the CUDA path is demonstrably free of this defect rather than merely different. In
jit/csrc/deepseek_v4/store.cuhthe page stride is declared 64-bit at every site:The page index there is narrow too (
const int32_t page = index >> kPageBits;), butint32_t * int64_tpromotes, so the multiply cannot wrap. The Triton port writes the structurally identicalpage * BYTES_PER_PAGE, except its stride is atl.constexprPython int that stays int32. This is a port that lost a type: in C++ the width is declared, in Triton it is implied by the value. That is consistent with no upstream report despite DeepSeek-V4 running widely on NVIDIA.DeepSeek-V4 family only. The single caller of the entry point is
srt/mem_cache/deepseek_v4_memory_pool.py; the kernels also hard-code MLA geometry (512 = 448 nope + 64 rope) and the 128-dim C4 indexer.A cache under ~2-4 GiB cannot produce a page index large enough to cross, so small deployments are unaffected regardless.
I grepped the rest of
python/sglang/kernelsfor the same shape (an index narrowed to int32 then multiplied by a page stride). The only other hit is alengthload indsv4_attn_metadata_kernels.pywhose multiply is unrelated, so this does not look like a widespread idiom.That audit searched for the pattern, though, not for Triton ports that dropped an
int64_ttheir CUDA original declared. If reviewers know of other such ports, they are worth the same look.Modifications
Widen
pageto int64 before the multiply, in both store kernels.slotis bounded byPAGE_SIZEand stays int32, so the added cost is one conversion per token per tile.test/registered/kernels/ops/kvcache/test_triton_store_cache_large_cache.pycovers both kernels at a slot past the crossing. The fixtures need a cache large enough for the crossing to exist: 4.00 GiB for the flashmla rope tile (bf16 view, stridebytes_per_page // 2) and 8.00 GiB for the indexer scales (f32 view,// 4); both skip when the device lacks the memory.The indexer kernel carries the identical narrowing. It was not observed to fault in this workload because its stride puts the crossing elsewhere, but the defect is the same and is fixed alongside.
Checklist
hipErrorIllegalAddresson the unpatched tree, all 3 pass on the patched one (4x MI355X)kernels/for the same narrowing: no other instance found🤖 Generated with Claude Code
CI States
Latest PR Test (Base): ❌ Run #36167697019
Latest PR Test (Extra): ❌ Run #36167696956
Latest PR Test (AMD ROCm 10): ❌ Run #36167697059