diff --git a/python/sglang/kernels/ops/attention/flash_mla_sm120.py b/python/sglang/kernels/ops/attention/flash_mla_sm120.py index 2193411a7..ec43c5508 100644 --- a/python/sglang/kernels/ops/attention/flash_mla_sm120.py +++ b/python/sglang/kernels/ops/attention/flash_mla_sm120.py @@ -13,6 +13,7 @@ separate region at the end of each page. import logging import math +from typing import Optional import torch import triton @@ -305,8 +306,15 @@ def _page_split_kernel( DST_SCALE_OFF: tl.constexpr, # 64 * 576 = 36864 RATIO: tl.constexpr, # 4 BLOCK_SIZE: tl.constexpr, + mask_ptr, + HAS_MASK: tl.constexpr, ): - """Fused page-split: copy data+scale for all sub-pages in one kernel.""" + """Fused page-split: copy data+scale for all sub-pages in one kernel. + + When HAS_MASK is set, only pages flagged in ``mask_ptr`` (int8, 1=touched) + are copied; untouched pages are skipped so the kernel no longer rewrites the + entire KV pool every decode step. + """ pid = tl.program_id(0) page_idx = pid // RATIO sub = pid % RATIO @@ -314,6 +322,10 @@ def _page_split_kernel( if page_idx >= N_pages: return + if HAS_MASK: + if tl.load(mask_ptr + page_idx) == 0: + return + src_base = src_ptr + page_idx * src_stride0 dst_base = dst_ptr + (page_idx * RATIO + sub) * dst_stride0 @@ -334,11 +346,43 @@ def _page_split_kernel( tl.store(dst_base + DST_SCALE_OFF + offs, vals, mask=mask) -def _split_kv_pages_to_64(kv_u8: torch.Tensor, src_pbs: int) -> torch.Tensor: +@triton.jit +def _page_mark_kernel( + indices_ptr, + mask_ptr, + N_idx, + SRC_PBS: tl.constexpr, + BLOCK: tl.constexpr, +): + """Mark touched source pages (1 byte each) from token-level indices. + + ``indices`` are token indices into the pbs=SRC_PBS SWA pool; -1 = invalid. + Each valid token marks ``mask[token // SRC_PBS] = 1``. Concurrent stores of + the same value 1 are safe (no atomic needed). + """ + pid = tl.program_id(0) + if pid >= N_idx: + return + idx = tl.load(indices_ptr + pid) + if idx < 0: + return + page = idx // SRC_PBS + tl.store(mask_ptr + page, 1) + + +def _split_kv_pages_to_64( + kv_u8: torch.Tensor, + src_pbs: int, + touched_indices: Optional[torch.Tensor] = None, +) -> torch.Tensor: """Split pbs=N footer-format pages into pbs=64 footer-format pages. - Uses a fused Triton kernel to do all sub-page copies in a single launch - instead of 8 separate copy kernels (4 sub-pages × 2 regions). + When ``touched_indices`` (token-level int32 indices into the pbs=src_pbs + SWA pool, -1 = invalid) is provided, only the source pages that actually + contain a referenced token are copied. This avoids rewriting the entire KV + pool on every decode step (only ~2*batch pages are touched vs the full + pool). The output buffer is persistent and reused across steps; untouched + dst pages simply retain their (unreferenced) stale data. """ assert src_pbs % _PBS_DST == 0 and src_pbs >= _PBS_DST if src_pbs == _PBS_DST: @@ -373,6 +417,35 @@ def _split_kv_pages_to_64(kv_u8: torch.Tensor, src_pbs: int) -> torch.Tensor: else: src_stride0 = src_2d.stride(0) + use_mask = touched_indices is not None and touched_indices.numel() > 0 + mask_ptr = src_2d # dummy, never dereferenced when HAS_MASK is False + if use_mask: + # Persistent per-device int8 mask, zeroed each call (cheap memset, + # captured cleanly by CUDA graph). 1 = page is referenced this step. + mkey = f"flash_mla_sm120_mask:{dev}" + mbuf = buffers.get(mkey) + if mbuf is None or mbuf.shape[0] < N: + # The first allocation can happen under inference mode (autotune), + # but the buffer is zeroed again later during CUDA graph capture + # outside inference mode -- an inference tensor cannot be mutated + # there, so force a normal tensor. + with torch.inference_mode(False): + mbuf = torch.empty(N, dtype=torch.int8, device=dev) + buffers[mkey] = mbuf + mask = mbuf[:N] + mask.zero_() + idx_flat = touched_indices.reshape(-1).contiguous() + if idx_flat.dtype != torch.int32: + idx_flat = idx_flat.to(torch.int32) + _page_mark_kernel[(idx_flat.numel(),)]( + idx_flat, + mask, + idx_flat.numel(), + src_pbs, # SRC_PBS + 1024, # BLOCK (unused, kept for JIT signature) + ) + mask_ptr = mask + grid = (N * ratio,) _page_split_kernel[grid]( src_2d, @@ -386,6 +459,8 @@ def _split_kv_pages_to_64(kv_u8: torch.Tensor, src_pbs: int) -> torch.Tensor: _PBS_DST * _NOPE_ROPE_STRIDE, # DST_SCALE_OFF = 36864 ratio, # RATIO = 4 1024, # BLOCK_SIZE + mask_ptr, + use_mask, # HAS_MASK ) bpt = _NOPE_ROPE_STRIDE + _SCALE_STRIDE # 584 @@ -424,10 +499,19 @@ def _flash_mla_flashinfer( B, _, H, D = q.shape # (batch, 1, num_heads, head_dim) dev = q.device + # Indices: no remapping needed (page-split preserves token addressing). + idx = indices.squeeze(1) if indices.dim() == 3 else indices + # --- Page-split: convert pbs=N kv_cache to pbs=64 view --- + # Only the SWA pages actually referenced by `idx` are copied (the rest of + # the persistent dst buffer is left untouched and never read). kv_u8 = k_cache.view(torch.uint8) if k_cache.dtype != torch.uint8 else k_cache src_pbs = k_cache.shape[1] if k_cache.ndim >= 3 else _PBS_SRC - kv_64 = _split_kv_pages_to_64(kv_u8, src_pbs) if src_pbs != _PBS_DST else kv_u8 + kv_64 = ( + _split_kv_pages_to_64(kv_u8, src_pbs, touched_indices=idx) + if src_pbs != _PBS_DST + else kv_u8 + ) extra_kv_u8 = ( extra_k_cache.view(torch.uint8) @@ -436,8 +520,6 @@ def _flash_mla_flashinfer( ) extra_kv_64 = extra_kv_u8 - # Indices: no remapping needed (page-split preserves token addressing). - idx = indices.squeeze(1) if indices.dim() == 3 else indices extra_idx = ( extra_indices.squeeze(1) if extra_indices is not None and extra_indices.dim() == 3 diff --git a/test/registered/kernels/ops/attention/test_flash_mla_backends.py b/test/registered/kernels/ops/attention/test_flash_mla_backends.py index 973dbf454..17867fb39 100644 --- a/test/registered/kernels/ops/attention/test_flash_mla_backends.py +++ b/test/registered/kernels/ops/attention/test_flash_mla_backends.py @@ -28,15 +28,20 @@ import torch from sglang.kernels.ops.attention import flash_mla_sm120 as fmod from sglang.kernels.ops.attention.flash_mla_sm120 import ( + _BYTES_PER_DST_PAGE, + _BYTES_PER_DST_PAGE_PADDED, _D, _NOPE_DIM, _NOPE_ROPE_STRIDE, _NUM_TILES, + _PBS_DST, + _PBS_SRC, _ROPE_DIM, _SCALE_STRIDE, _TILE_SIZE, _gather_and_dequant, _sm120_sparse_decode_fwd, + _split_kv_pages_to_64, flash_mla_with_kvcache_sm120, ) from sglang.kernels.ops.attention.flash_mla_sm120_triton import ( @@ -44,6 +49,7 @@ from sglang.kernels.ops.attention.flash_mla_sm120_triton import ( _merge_partial_attn, flash_mla_sparse_decode_triton, ) +from sglang.srt.runtime_context import get_resources from sglang.test.ci.ci_register import register_cuda_ci from sglang.test.test_utils import CustomTestCase @@ -499,6 +505,130 @@ class TestEntryPointDispatch(CustomTestCase): ) +@unittest.skipUnless(_IS_SM120, "SM120 (compute capability 12.0) required") +class TestTouchedPageSplit(CustomTestCase): + """The pbs=256 -> pbs=64 split must rewrite only referenced source pages. + + The destination buffer is persistent across decode steps, so the masked + split is only correct if it (a) copies every byte of the data and scale + regions of a marked page's sub-pages and (b) leaves everything else alone + -- the per-sub-page alignment tail and all sub-pages of an unmarked source + page. A widened copy or a dropped mask check violates the persistent-buffer + contract, which checking only the copied regions would not catch. + """ + + @classmethod + def setUpClass(cls): + if not torch.cuda.is_available(): + raise unittest.SkipTest("CUDA required") + cls.device = torch.device("cuda") + + def test_only_marked_pages_are_split(self): + num_pages = 3 + ratio = _PBS_SRC // _PBS_DST + sentinel = 0xA5 + + k_cache, _ = _build_kvcache(num_pages, _PBS_SRC, device=self.device, seed=17) + # Production takes the raw 2D byte view of the (N, pbs, 1, bpt) cache. + src_stride0 = k_cache.stride(0) + src_2d = torch.as_strided( + k_cache.view(torch.uint8), (num_pages, src_stride0), (src_stride0, 1) + ) + + # Seed the exact persistent buffers production reuses, so the sentinel + # bytes below are the ones the kernel writes into. + dev = k_cache.device + buffers = get_resources().buffers + split_key = f"flash_mla_sm120_split:{dev}" + mask_key = f"flash_mla_sm120_mask:{dev}" + missing = object() + for key in (split_key, mask_key): + old = buffers.get(key, missing) + + def _restore(key=key, old=old): + if old is missing: + buffers.pop(key, None) + else: + buffers[key] = old + + self.addCleanup(_restore) + + dst = torch.full( + (num_pages * ratio, _BYTES_PER_DST_PAGE_PADDED), + sentinel, + dtype=torch.uint8, + device=self.device, + ) + buffers[split_key] = dst + # The mask is deliberately not preseeded: production allocates it on the + # first call, which happens during inference-mode autotune. + + # Two tokens in source page 0, one in page 2, one invalid; page 1 idle. + token_ids = torch.tensor( + [0, 5, 2 * _PBS_SRC + 3, -1], dtype=torch.int32, device=self.device + ) + + # Autotune-like first call: allocates the persistent mask under + # inference mode. It must still be a normal tensor, or the later + # zero_() (CUDA graph capture, outside inference mode) would fail. + with torch.inference_mode(): + _split_kv_pages_to_64( + k_cache.view(torch.uint8), _PBS_SRC, touched_indices=token_ids + ) + torch.cuda.synchronize() + self.assertFalse( + buffers[mask_key].is_inference(), + "persistent mask must not be an inference tensor", + ) + + # Restore the pre-step state: sentinel destination and nonzero mask + # storage that the call must zero before marking. + dst.fill_(sentinel) + buffers[mask_key].fill_(-7) + + out_pages = _split_kv_pages_to_64( + k_cache.view(torch.uint8), _PBS_SRC, touched_indices=token_ids + ) + torch.cuda.synchronize() + + self.assertEqual( + out_pages.shape, (num_pages * ratio, _PBS_DST, 1, _BYTES_PER_TOKEN) + ) + self.assertEqual(buffers[mask_key].tolist(), [1, 0, 1]) + + data_per_sub = _PBS_DST * _NOPE_ROPE_STRIDE + scale_per_sub = _PBS_DST * _SCALE_STRIDE + src_scale_off = _PBS_SRC * _NOPE_ROPE_STRIDE + for page in (0, 2): + for sub in range(ratio): + dst_page = page * ratio + sub + out = dst[dst_page] + torch.testing.assert_close( + out[:data_per_sub], + src_2d[page, sub * data_per_sub : (sub + 1) * data_per_sub], + atol=0, + rtol=0, + msg=f"data region mismatch for dst page {dst_page}", + ) + scale_off = src_scale_off + sub * scale_per_sub + torch.testing.assert_close( + out[data_per_sub:_BYTES_PER_DST_PAGE], + src_2d[page, scale_off : scale_off + scale_per_sub], + atol=0, + rtol=0, + msg=f"scale region mismatch for dst page {dst_page}", + ) + self.assertTrue( + bool((out[_BYTES_PER_DST_PAGE:] == sentinel).all()), + f"alignment padding of dst page {dst_page} was overwritten", + ) + + self.assertTrue( + bool((dst[ratio : 2 * ratio] == sentinel).all()), + "sub-pages of untouched source page 1 were rewritten", + ) + + if __name__ == "__main__": import sys