Skip to content

[Kernel] Fix int32 page-offset overflow in the paged KV store kernels - #40846

Closed
tianxiaojiang4 wants to merge 1 commit into
sgl-project:mainfrom
tianxiaojiang4:fix/dsv4-kv-store-int32-page-offset
Closed

tianxiaojiang4 wants to merge 1 commit into
sgl-project:mainfrom
tianxiaojiang4:fix/dsv4-kv-store-int32-page-offset

Conversation

@tianxiaojiang4

@tianxiaojiang4 tianxiaojiang4 commented Sep 23, 2026 •

Copy link
Copy Markdown
Contributor

Motivation

_triton_fused_store_flashmla_kernel and _triton_fused_store_indexer_kernel narrow the KV slot index to int32 and then multiply it by the page stride:

loc  = tl.load(indices_ptr + token_id).to(tl.int32)
page = loc // PAGE_SIZE
rope_bf16_offset = page * BYTES_PER_PAGE_BF16 + slot * BF16_SLOT_ELEMS + ...
tl.store(cache_bf16_ptr + rope_bf16_offset, rope_vals)

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 is bytes_per_page // 2 = 74880 and int32 is exhausted at page 28679, i.e. slot index 7,341,824:

max slot index page page * 74880 result
6,950,144 27,149 2,032,917,120 healthy
7,340,743 28,674 2,147,109,120 (99.983% of 2^31) faulted in the next batch

AMD_SERIALIZE_KERNEL=3 attributes the fault to _triton_fused_store_flashmla_kernel on 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-cache slots 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:

  • concurrency (8, 16 and 64 all fail at the same total)
  • request count (64 requests of 121k tokens and 128 of 60.5k fail at the same total)
  • pool sizing (verified across a 4.8x SWA range and a 4.6x compressed-pool range; no pool exceeded 2% occupancy at the fault)

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:

concurrency uncached tokens/s time to cross
1-4 0.6k-1.9k 1-3 h, survives the run
8 2.0k-4.2k 29-60 min
32 7.0k-15.3k 8-17 min
64 14.3k-22.9k 5-9 min

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:

if is_hip_runtime():
    triton_fused_store_cache(...)        # these kernels
else:
    module = _jit_fused_store_module(...)  # CUDA uses a JIT kernel instead
  • ROCm only, and the CUDA path is demonstrably free of this defect rather than merely different. In jit/csrc/deepseek_v4/store.cuh the page stride is declared 64-bit at every site:

    static constexpr int64_t kPageBytes = deepseek_v4::kv_page_bytes<kLayout>(kPageSize);
    constexpr int64_t        kPageBytes = 132 << kPageBits;
    static constexpr int64_t kPageBytes = 132 * kPageSize;

    The page index there is narrow too (const int32_t page = index >> kPageBits;), but int32_t * int64_t promotes, so the multiply cannot wrap. The Triton port writes the structurally identical page * BYTES_PER_PAGE, except its stride is a tl.constexpr Python 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/kernels for the same shape (an index narrowed to int32 then multiplied by a page stride). The only other hit is a length load in dsv4_attn_metadata_kernels.py whose 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_t their CUDA original declared. If reviewers know of other such ports, they are worth the same look.

Modifications

Widen page to int64 before the multiply, in both store kernels. slot is bounded by PAGE_SIZE and stays int32, so the added cost is one conversion per token per tile.

test/registered/kernels/ops/kvcache/test_triton_store_cache_large_cache.py covers 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, stride bytes_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

  • Root cause confirmed by direct observation, not inference: an instrumented build prints the computed offsets and the last healthy sample sits 1,081 slots below the int32 crossing
  • Faulting kernel confirmed under serialized dispatch on all ranks
  • Fix verification: the patched build runs to 11.6M tokens and slot 10,987,776, 1.5x past the crossing, with 0 faults
  • Regression test fails before the fix and passes after: all 3 cases error with hipErrorIllegalAddress on the unpatched tree, all 3 pass on the patched one (4x MI355X)
  • Audited the rest of 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

@tianxiaojiang4

Copy link
Copy Markdown
Contributor Author

Fix verified

Same 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.

unpatched patched
max slot index reached 7,340,743 10,987,776
page * BYTES_PER_PAGE_BF16 reached 2,147,109,120 (99.98% of 2^31) 3,213,924,480 (1.50x of 2^31)
cumulative prefilled tokens 7.744M, faulted 11.616M, healthy
hipErrorIllegalAddress occurrences 47-152 across runs 0

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:

  • Fix verification run: passes, 1.5x past the crossing
  • Reviewer check still requested: whether any other paged-store kernel shares this narrowing

🤖 Generated with Claude Code

@tianxiaojiang4

Copy link
Copy Markdown
Contributor Author

Regression test verified against both trees

The 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.

ARM=unfixed   int64 sites in tree: 0
  FAILED test_store_past_int32_page_offset[triton_fused_store_flashmla-512-2-149760]
  FAILED test_store_past_int32_page_offset[triton_fused_store_indexer-128-4-33792]
  FAILED test_flashmla_high_slot_does_not_clobber_other_pages
  3 failed  (all: torch.AcceleratorError: CUDA error: an illegal memory access was encountered)

ARM=fixed     int64 sites in tree: 2
  3 passed

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

@tianxiaojiang4
tianxiaojiang4 marked this pull request as ready for review September 23, 2026 17:39
@tianxiaojiang4
tianxiaojiang4 force-pushed the fix/dsv4-kv-store-int32-page-offset branch from 2604c2f to 753872f Compare September 23, 2026 19:04
@tianxiaojiang4

Copy link
Copy Markdown
Contributor Author

@kkHuang-amd @hnyls2002 @amd-danli103 — tagging you from git log on the touched paths; apologies if this misses the mark.

This fixes an int32 overflow of the page offset in the two ROCm Triton KV store kernels. page * BYTES_PER_PAGE wraps once the paged cache passes ~2 GiB, so the store scatters outside the buffer and the server dies with hipErrorIllegalAddress. DeepSeek-V4/V4.1 on ROCm only. The CUDA path is unaffected: store.cuh already declares kPageBytes as int64_t, so the same multiply promotes there.

Why each of you:

Measured on MI355X with AMD_SERIALIZE_KERNEL=3: unpatched dies at slot 7,340,743, an offset at 99.98% of 2^31; patched reaches slot 10,987,776, 1.5x past the crossing, with no faults. The added regression test fails without the fix and passes with it.

One ask: could someone with permission add the run-ci label (/tag-and-rerun-ci)? I am not in CI_PERMISSIONS.json, so the test suites have not run on this yet.

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>
@tianxiaojiang4
tianxiaojiang4 force-pushed the fix/dsv4-kv-store-int32-page-offset branch from 753872f to 200e6b7 Compare September 25, 2026 17:31
@tianxiaojiang4

Copy link
Copy Markdown
Contributor Author

Closed in favor of #41159

Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Projects

None yet

Development

Successfully merging this pull request may close these issues.

1 participant