[mxfp8-kv] Skip writes to the reserved CUDA-graph padding slot (#35351)

Co-authored-by: Claude Opus 4.8 <noreply@anthropic.com>
Co-authored-by: Ke Bao <ispobaoke@gmail.com>
Co-authored-by: Sam Shleifer <sam@thinkingmachines.ai>
This commit is contained in:
Sam Shleifer
2026-09-22 01:21:25 +08:00
committed by GitHub
co-authored by Claude Opus 4.8 Ke Bao Sam Shleifer
parent 008470abd8
commit 5a6a1bb883
4 changed files with 165 additions and 15 deletions
@@ -40,8 +40,9 @@ def _store_sf_interleaved_kernel(
tok_offsets = tok_start + tl.arange(0, BLOCK_T)
mask = tok_offsets < num_tokens
# Load slot indices
# Slot 0 is the reserved CUDA-graph padding sink; skip writes to it.
slots = tl.load(loc_ptr + tok_offsets, mask=mask, other=0)
mask = mask & (slots != 0)
page_offsets = slots % page_size
page_idxs = slots // page_size
@@ -143,6 +143,9 @@ def _mxfp8_quant_store_qkv_kernel(
tl.store(sfq_ptr + (t * NQ + r) * SF + blk, sf)
else:
myloc = tl.load(loc_ptr + t).to(tl.int64)
# Slot 0 is the reserved CUDA-graph padding sink; skip writes to it.
if myloc == 0:
return
if r < NQ + NKV:
h = r - NQ
cache = kc_ptr
+22 -14
View File
@@ -3650,20 +3650,28 @@ class MHATokenToKVPoolMXFP8(MHATokenToKVPool):
)
return
from sglang.srt.model_executor.runner import get_is_capture_mode
if get_is_capture_mode() and self.alt_stream is not None:
current_stream = self.device_module.current_stream()
self.alt_stream.wait_stream(current_stream)
self.k_buffer[idx][loc] = cache_k
self._write_scales(idx, loc, k_scale, v_scale)
with self.device_module.stream(self.alt_stream):
self.v_buffer[idx][loc] = cache_v
current_stream.wait_stream(self.alt_stream)
else:
self.k_buffer[idx][loc] = cache_k
self.v_buffer[idx][loc] = cache_v
self._write_scales(idx, loc, k_scale, v_scale)
# store_cache and store_sf_interleaved skip the reserved CUDA-graph
# padding slot 0 in-kernel, matching the bf16 pool.
row_bytes = self.head_num * self.head_dim * self.store_dtype.itemsize
v_row_bytes = self.head_num * self.v_head_dim * self.store_dtype.itemsize
assert _is_cuda and can_use_store_cache(row_bytes, v_row_bytes), (
f"MXFP8 KV cache requires CUDA and store_cache-compatible rows, "
f"got _is_cuda={_is_cuda}, {row_bytes=}, {v_row_bytes=}"
)
assert self.mxfp8_sf_interleaved, (
"MXFP8 KV cache requires the page_size=128 interleaved scale layout"
)
store_cache(
cache_k.reshape(loc.shape[0], -1),
cache_v.reshape(loc.shape[0], -1),
self.k_buffer[idx].view(-1, row_bytes // self.store_dtype.itemsize),
self.v_buffer[idx].view(-1, v_row_bytes // self.store_dtype.itemsize),
loc,
row_bytes=row_bytes,
v_row_bytes=v_row_bytes,
size_limit=self.size + self.page_size,
)
self._write_scales(idx, loc, k_scale, v_scale)
def _write_scales(self, idx, loc, k_scale, v_scale):
"""Write per-token UE8M0 K/V scales — interleaved into the FA4
@@ -0,0 +1,138 @@
"""MXFP8 KV cache must never write the reserved CUDA-graph padding slot.
Padding lanes carry undefined activations that quantize to NaN payload and
0xFF e8m0 scales; attention reads slot 0 back for padded page-table entries,
so a poisoned slot 0 defeats probability masking (0 * NaN = NaN in PV).
Asserts require slot 0 to stay exactly zero, so finite-garbage writes fail too.
"""
import pytest
import torch
from sglang.srt.utils import get_device_sm
from sglang.test.ci.ci_register import register_cuda_ci
register_cuda_ci(est_time=60, stage="base-b-kernel-unit", runner_config="4-gpu-b200")
requires_sm100 = pytest.mark.skipif(
not torch.cuda.is_available() or get_device_sm() < 100,
reason="MXFP8 KV cache requires SM100+",
)
DEV, HD, PS, NHKV = "cuda", 128, 128, 2
def _make_pool(**kwargs):
from sglang.srt.mem_cache.memory_pool import MHATokenToKVPoolMXFP8
return MHATokenToKVPoolMXFP8(
size=4 * PS,
page_size=PS,
dtype=torch.float8_e4m3fn,
head_num=NHKV,
head_dim=HD,
layer_num=1,
device=DEV,
enable_memory_saver=False,
**kwargs,
)
class _Layer:
layer_id = 0
def _quantize(k, v):
from sglang.kernels.ops.quantization.mxfp8_quant import to_mxfp8
km, vm = to_mxfp8(k), to_mxfp8(v)
return (
km.data,
vm.data,
km.scale.view(torch.float8_e8m0fnu),
vm.scale.view(torch.float8_e8m0fnu),
)
def _assert_slot0_zero(pool):
kc, vc = pool.get_kv_buffer(0)
ksf, vsf = pool.get_kv_scale_buffer(0)
s0k = kc.view(-1, PS, NHKV, HD)[0, 0].view(torch.uint8)
s0v = vc.view(-1, PS, NHKV, HD)[0, 0].view(torch.uint8)
assert int(s0k.sum()) == 0, "reserved slot K payload written"
assert int(s0v.sum()) == 0, "reserved slot V payload written"
zero_loc = torch.zeros(1, dtype=torch.int64, device=DEV)
s0_ksf = pool._read_sf_interleaved(ksf, zero_loc).view(torch.uint8)
s0_vsf = pool._read_sf_interleaved(vsf, zero_loc).view(torch.uint8)
assert int(s0_ksf.sum()) == 0, "reserved slot K scales written"
assert int(s0_vsf.sum()) == 0, "reserved slot V scales written"
@requires_sm100
def test_direct_path_skips_reserved_slot():
torch.manual_seed(0)
pool = _make_pool()
k = torch.randn(4, NHKV, HD, dtype=torch.bfloat16, device=DEV) * 0.5
v = torch.randn(4, NHKV, HD, dtype=torch.bfloat16, device=DEV) * 0.5
k[[0, 2]] = float("nan")
v[[0, 2]] = float("nan")
kq, vq, ks, vs = _quantize(k, v)
loc = torch.tensor([0, 7, 0, 9], dtype=torch.int64, device=DEV)
pool.set_kv_buffer(_Layer(), loc, kq, vq, ks, vs)
_assert_slot0_zero(pool)
kc, _ = pool.get_kv_buffer(0)
got = kc.view(-1, PS, NHKV, HD)[0, 7].view(torch.uint8)
assert torch.equal(got, kq[1].view(torch.uint8)), "non-reserved write corrupted"
@requires_sm100
def test_fused_quant_store_path_skips_reserved_slot():
"""k_scale=None routes to the fused quant_store_kv_mxfp8 kernel."""
torch.manual_seed(2)
pool = _make_pool()
k = torch.randn(4, NHKV, HD, dtype=torch.bfloat16, device=DEV) * 0.5
v = torch.randn(4, NHKV, HD, dtype=torch.bfloat16, device=DEV) * 0.5
k[[1, 3]] = float("nan")
v[[1, 3]] = float("nan")
loc = torch.tensor([5, 0, 9, 0], dtype=torch.int64, device=DEV)
pool.set_kv_buffer(_Layer(), loc, k, v)
_assert_slot0_zero(pool)
kc, _ = pool.get_kv_buffer(0)
valid = kc.view(-1, PS, NHKV, HD)[0, 5].float()
assert not torch.isnan(valid).any() and valid.abs().sum() > 0, (
"fused valid write lost"
)
@requires_sm100
def test_set_kv_buffer_is_cuda_graph_capture_safe():
torch.manual_seed(5)
pool = _make_pool(enable_alt_stream=False)
k = torch.randn(2, NHKV, HD, dtype=torch.bfloat16, device=DEV) * 0.5
v = torch.randn(2, NHKV, HD, dtype=torch.bfloat16, device=DEV) * 0.5
kq, vq, ks, vs = _quantize(k, v)
loc = torch.tensor([0, 7], dtype=torch.int64, device=DEV)
for _ in range(2): # warmup
pool.set_kv_buffer(_Layer(), loc, kq, vq, ks, vs)
torch.cuda.synchronize()
g = torch.cuda.CUDAGraph()
with torch.cuda.graph(g):
pool.set_kv_buffer(_Layer(), loc, kq, vq, ks, vs)
g.replay()
torch.cuda.synchronize()
_assert_slot0_zero(pool)
kc, _ = pool.get_kv_buffer(0)
got = kc.view(-1, PS, NHKV, HD)[0, 7].view(torch.uint8)
assert torch.equal(got, kq[1].view(torch.uint8)), "captured valid write lost"
if __name__ == "__main__":
import sys
sys.exit(pytest.main([__file__, "-v"]))