From 5a6a1bb88354b4d0ed13c0424cb6e3bc7539c674 Mon Sep 17 00:00:00 2001 From: Sam Shleifer Date: Mon, 21 Sep 2026 13:21:25 -0400 Subject: [PATCH] [mxfp8-kv] Skip writes to the reserved CUDA-graph padding slot (#35351) Co-authored-by: Claude Opus 4.8 Co-authored-by: Ke Bao Co-authored-by: Sam Shleifer --- .../ops/quantization/mxfp8_interleave_sf.py | 3 +- .../kernels/ops/quantization/mxfp8_quant.py | 3 + python/sglang/srt/mem_cache/memory_pool.py | 36 +++-- .../test_mxfp8_kv_reserved_slot.py | 138 ++++++++++++++++++ 4 files changed, 165 insertions(+), 15 deletions(-) create mode 100644 test/registered/kernel/quantization/test_mxfp8_kv_reserved_slot.py diff --git a/python/sglang/kernels/ops/quantization/mxfp8_interleave_sf.py b/python/sglang/kernels/ops/quantization/mxfp8_interleave_sf.py index 8a9ce307d..f44609b52 100644 --- a/python/sglang/kernels/ops/quantization/mxfp8_interleave_sf.py +++ b/python/sglang/kernels/ops/quantization/mxfp8_interleave_sf.py @@ -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 diff --git a/python/sglang/kernels/ops/quantization/mxfp8_quant.py b/python/sglang/kernels/ops/quantization/mxfp8_quant.py index 6d5608527..a9a59b545 100644 --- a/python/sglang/kernels/ops/quantization/mxfp8_quant.py +++ b/python/sglang/kernels/ops/quantization/mxfp8_quant.py @@ -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 diff --git a/python/sglang/srt/mem_cache/memory_pool.py b/python/sglang/srt/mem_cache/memory_pool.py index 3df1bfad7..0ae0ec0af 100644 --- a/python/sglang/srt/mem_cache/memory_pool.py +++ b/python/sglang/srt/mem_cache/memory_pool.py @@ -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 diff --git a/test/registered/kernel/quantization/test_mxfp8_kv_reserved_slot.py b/test/registered/kernel/quantization/test_mxfp8_kv_reserved_slot.py new file mode 100644 index 000000000..8337e2a67 --- /dev/null +++ b/test/registered/kernel/quantization/test_mxfp8_kv_reserved_slot.py @@ -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"]))