[SM120] Add FlashInfer sparse MLA decode for DSv4-Flash (#27455)

Co-authored-by: AliceChenyy <alicechenyy@users.noreply.github.com>
This commit is contained in:
eeecho
2026-06-29 16:27:16 -07:00
committed by GitHub
co-authored by AliceChenyy
parent b8c25bfaa7
commit f0bf96390b
3 changed files with 261 additions and 7 deletions
+2
View File
@@ -582,6 +582,8 @@ class Envs:
# None = standard attention. See https://arxiv.org/abs/2512.12087 # None = standard attention. See https://arxiv.org/abs/2512.12087
SGLANG_SKIP_SOFTMAX_PREFILL_THRESHOLD_SCALE_FACTOR = EnvFloat(None) SGLANG_SKIP_SOFTMAX_PREFILL_THRESHOLD_SCALE_FACTOR = EnvFloat(None)
SGLANG_SKIP_SOFTMAX_DECODE_THRESHOLD_SCALE_FACTOR = EnvFloat(None) SGLANG_SKIP_SOFTMAX_DECODE_THRESHOLD_SCALE_FACTOR = EnvFloat(None)
# SM120 FlashMLA decode backend: "flashinfer" (default), "triton", or "torch".
SGLANG_SM120_FLASHMLA_BACKEND = EnvStr("flashinfer")
# Triton # Triton
SGLANG_TRITON_DECODE_ATTN_STATIC_KV_SPLITS = EnvBool(False) SGLANG_TRITON_DECODE_ATTN_STATIC_KV_SPLITS = EnvBool(False)
@@ -12,9 +12,13 @@ separate region at the end of each page.
""" """
import logging import logging
import os import math
import torch import torch
import triton
import triton.language as tl
from sglang.srt.environ import envs
logger = logging.getLogger(__name__) logger = logging.getLogger(__name__)
@@ -193,17 +197,15 @@ def _sm120_sparse_decode_fwd(
return out.to(torch.bfloat16), lse.permute(0, 2, 1) return out.to(torch.bfloat16), lse.permute(0, 2, 1)
# Default SM120 FlashMLA backend: "triton" (optimized) or "torch" (pure-PyTorch fallback). # SM120 FlashMLA: default FlashInfer (CUTLASS SM120 sparse MLA decode).
# Controlled by SGLANG_SM120_TRITON_FLASHMLA env var (1=triton, 0=torch). # Override with SGLANG_SM120_FLASHMLA_BACKEND=triton|torch to force fallback.
_sm120_default_backend = ( _sm120_default_backend = envs.SGLANG_SM120_FLASHMLA_BACKEND.get()
"triton" if os.environ.get("SGLANG_SM120_TRITON_FLASHMLA", "1") == "1" else "torch"
)
def flash_mla_with_kvcache_sm120(**kwargs): def flash_mla_with_kvcache_sm120(**kwargs):
"""SM120 FlashMLA sparse decode entry point. """SM120 FlashMLA sparse decode entry point.
Dispatches to the Triton kernel (default) or PyTorch fallback. Dispatches to FlashInfer (default if available), Triton, or PyTorch fallback.
""" """
q = kwargs["q"] q = kwargs["q"]
k_cache = kwargs["k_cache"] k_cache = kwargs["k_cache"]
@@ -218,6 +220,20 @@ def flash_mla_with_kvcache_sm120(**kwargs):
extra_indices = kwargs.get("extra_indices_in_kvcache") extra_indices = kwargs.get("extra_indices_in_kvcache")
extra_topk_length = kwargs.get("extra_topk_length") extra_topk_length = kwargs.get("extra_topk_length")
if _sm120_default_backend == "flashinfer":
return _flash_mla_flashinfer(
q,
k_cache,
indices,
topk_length,
attn_sink,
head_dim_v,
softmax_scale,
extra_k_cache,
extra_indices,
extra_topk_length,
)
if _sm120_default_backend == "triton": if _sm120_default_backend == "triton":
from sglang.srt.layers.attention.flash_mla_sm120_triton import ( from sglang.srt.layers.attention.flash_mla_sm120_triton import (
flash_mla_sparse_decode_triton, flash_mla_sparse_decode_triton,
@@ -250,3 +266,199 @@ def flash_mla_with_kvcache_sm120(**kwargs):
extra_topk_length, extra_topk_length,
) )
return (out, lse) return (out, lse)
# --- Page-split utilities: pbs=256 → pbs=64 ---
# SGLang SWA KV cache footer layout per 256-token page:
# [data: 256 * 576 bytes] [scale: 256 * 8 bytes] [padding]
# FlashInfer decode_dsv4 expects per 64-token page:
# [data: 64 * 576 bytes] [scale: 64 * 8 bytes] [padding to 37440]
_PBS_SRC = 256 # SGLang physical page size
_PBS_DST = 64 # FlashInfer page_block_size
_NOPE_ROPE_STRIDE = 576 # bytes per token for nope+rope
_SCALE_STRIDE = 8 # bytes per token for scale (7 + 1 pad)
_BYTES_PER_DST_PAGE = (
_PBS_DST * _NOPE_ROPE_STRIDE + _PBS_DST * _SCALE_STRIDE
) # 64*576 + 64*8 = 37376 + 512 = 37888
# Padded to 576 alignment
_BYTES_PER_DST_PAGE_PADDED = math.ceil(_BYTES_PER_DST_PAGE / 576) * 576 # 37440
# Pre-allocated buffer for page-split output per device (lazily sized).
_split_buf = {} # device -> tensor
@triton.jit
def _page_split_kernel(
src_ptr,
dst_ptr,
N_pages,
src_stride0: tl.constexpr,
dst_stride0: tl.constexpr,
DATA_PER_SUB: tl.constexpr, # 64 * 576 = 36864
SCALE_PER_SUB: tl.constexpr, # 64 * 8 = 512
SRC_SCALE_OFF: tl.constexpr, # 256 * 576 = 147456
DST_SCALE_OFF: tl.constexpr, # 64 * 576 = 36864
RATIO: tl.constexpr, # 4
BLOCK_SIZE: tl.constexpr,
):
"""Fused page-split: copy data+scale for all sub-pages in one kernel."""
pid = tl.program_id(0)
page_idx = pid // RATIO
sub = pid % RATIO
if page_idx >= N_pages:
return
src_base = src_ptr + page_idx * src_stride0
dst_base = dst_ptr + (page_idx * RATIO + sub) * dst_stride0
# Copy data region: DATA_PER_SUB bytes from src offset sub*DATA_PER_SUB
data_src_off = sub * DATA_PER_SUB
for start in tl.range(0, DATA_PER_SUB, BLOCK_SIZE):
offs = start + tl.arange(0, BLOCK_SIZE)
mask = offs < DATA_PER_SUB
vals = tl.load(src_base + data_src_off + offs, mask=mask)
tl.store(dst_base + offs, vals, mask=mask)
# Copy scale region: SCALE_PER_SUB bytes
scale_src_off = SRC_SCALE_OFF + sub * SCALE_PER_SUB
for start in tl.range(0, SCALE_PER_SUB, BLOCK_SIZE):
offs = start + tl.arange(0, BLOCK_SIZE)
mask = offs < SCALE_PER_SUB
vals = tl.load(src_base + scale_src_off + offs, mask=mask)
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:
"""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).
"""
assert src_pbs % _PBS_DST == 0 and src_pbs >= _PBS_DST
if src_pbs == _PBS_DST:
return kv_u8
N = kv_u8.shape[0]
ratio = src_pbs // _PBS_DST
num_dst_pages = N * ratio
dev = kv_u8.device
buf = _split_buf.get(dev)
if buf is None or buf.shape[0] < num_dst_pages:
buf = torch.empty(
num_dst_pages,
_BYTES_PER_DST_PAGE_PADDED,
dtype=torch.uint8,
device=dev,
)
_split_buf[dev] = buf
out = buf[:num_dst_pages]
# Get raw 2D view of source
src_2d = kv_u8
if src_2d.ndim == 4:
src_stride0 = src_2d.stride(0)
src_2d = torch.as_strided(src_2d, (N, src_stride0), (src_stride0, 1))
else:
src_stride0 = src_2d.stride(0)
grid = (N * ratio,)
_page_split_kernel[grid](
src_2d,
out,
N,
src_stride0,
_BYTES_PER_DST_PAGE_PADDED,
_PBS_DST * _NOPE_ROPE_STRIDE, # DATA_PER_SUB = 36864
_PBS_DST * _SCALE_STRIDE, # SCALE_PER_SUB = 512
src_pbs * _NOPE_ROPE_STRIDE, # SRC_SCALE_OFF = 147456
_PBS_DST * _NOPE_ROPE_STRIDE, # DST_SCALE_OFF = 36864
ratio, # RATIO = 4
1024, # BLOCK_SIZE
)
bpt = _NOPE_ROPE_STRIDE + _SCALE_STRIDE # 584
return out.as_strided(
(num_dst_pages, _PBS_DST, 1, bpt),
(_BYTES_PER_DST_PAGE_PADDED, bpt, bpt, 1),
)
def _flash_mla_flashinfer(
q,
k_cache,
indices,
topk_length,
attn_sink,
head_dim_v,
softmax_scale,
extra_k_cache,
extra_indices,
extra_topk_length,
):
"""FlashInfer SM120 sparse MLA via sparse_mla_sm120_decode_dsv4.
SGLang SWA pool uses page_size=256 (footer format: 256*576 bytes data + 256*8 bytes scale).
FlashInfer decode_dsv4 fast path requires page_block_size=64 (footer: 64*576 + 64*8).
We split 256-token pages into 4 virtual 64-token pages.
Token indices are invariant under page-split (identity mapping).
"""
from flashinfer.mla._sparse_mla_sm120 import sparse_mla_sm120_decode_dsv4
B, _, H, D = q.shape # (batch, 1, num_heads, head_dim)
dev = q.device
# --- Page-split: convert pbs=N kv_cache to pbs=64 view ---
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
extra_kv_u8 = (
extra_k_cache.view(torch.uint8)
if extra_k_cache is not None and extra_k_cache.dtype != torch.uint8
else extra_k_cache
)
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
else extra_indices
)
output = torch.empty(B, H, head_dim_v, dtype=torch.bfloat16, device=dev)
out_lse = torch.empty(B, H, dtype=torch.float32, device=dev)
# Pre-allocate split-K scratch for decode-dsv4 fast path.
topk = idx.shape[-1]
extra_topk = extra_idx.shape[-1] if extra_idx is not None else 0
_BI = 64
num_splits = (topk + _BI - 1) // _BI + (
(extra_topk + _BI - 1) // _BI if extra_topk > 0 else 0
)
mid_out = torch.empty(
B, H, num_splits, head_dim_v, dtype=torch.bfloat16, device=dev
)
mid_lse = torch.empty(B, H, num_splits, dtype=torch.float32, device=dev)
sparse_mla_sm120_decode_dsv4(
q=q.squeeze(1) if q.ndim == 4 else q,
kv_cache=kv_64,
indices=idx,
mid_out=mid_out,
mid_lse=mid_lse,
output=output,
out_lse=out_lse,
sm_scale=softmax_scale,
topk_length=topk_length,
attn_sink=attn_sink,
extra_kv_cache=extra_kv_64,
extra_indices=extra_idx,
extra_topk_length=extra_topk_length,
)
return (output.unsqueeze(1), None)
@@ -53,6 +53,8 @@ register_cuda_ci(est_time=45, stage="base-b", runner_config="1-gpu-large")
# Per-token byte layout # Per-token byte layout
_BYTES_PER_TOKEN = _NOPE_ROPE_STRIDE + _SCALE_STRIDE # 576 + 8 = 584 _BYTES_PER_TOKEN = _NOPE_ROPE_STRIDE + _SCALE_STRIDE # 576 + 8 = 584
_IS_SM120 = torch.cuda.is_available() and torch.cuda.get_device_capability() == (12, 0)
def _build_kvcache( def _build_kvcache(
num_pages: int, num_pages: int,
@@ -179,6 +181,7 @@ def _build_q_indices(
return q, indices return q, indices
@unittest.skipUnless(_IS_SM120, "SM120 (compute capability 12.0) required")
class TestGatherAndDequant(CustomTestCase): class TestGatherAndDequant(CustomTestCase):
@classmethod @classmethod
def setUpClass(cls): def setUpClass(cls):
@@ -228,6 +231,7 @@ class TestGatherAndDequant(CustomTestCase):
) )
@unittest.skipUnless(_IS_SM120, "SM120 (compute capability 12.0) required")
class TestSparseDecodeTritonVsTorch(CustomTestCase): class TestSparseDecodeTritonVsTorch(CustomTestCase):
@classmethod @classmethod
def setUpClass(cls): def setUpClass(cls):
@@ -340,6 +344,7 @@ class TestSparseDecodeTritonVsTorch(CustomTestCase):
) )
@unittest.skipUnless(_IS_SM120, "SM120 (compute capability 12.0) required")
class TestApplyAttnSink(CustomTestCase): class TestApplyAttnSink(CustomTestCase):
@classmethod @classmethod
def setUpClass(cls): def setUpClass(cls):
@@ -382,6 +387,7 @@ class TestApplyAttnSink(CustomTestCase):
) )
@unittest.skipUnless(_IS_SM120, "SM120 (compute capability 12.0) required")
class TestMergePartialAttn(CustomTestCase): class TestMergePartialAttn(CustomTestCase):
@classmethod @classmethod
def setUpClass(cls): def setUpClass(cls):
@@ -423,6 +429,7 @@ class TestMergePartialAttn(CustomTestCase):
torch.testing.assert_close(merged_lse, lse1, atol=1e-5, rtol=1e-5) torch.testing.assert_close(merged_lse, lse1, atol=1e-5, rtol=1e-5)
@unittest.skipUnless(_IS_SM120, "SM120 (compute capability 12.0) required")
class TestEntryPointDispatch(CustomTestCase): class TestEntryPointDispatch(CustomTestCase):
@classmethod @classmethod
def setUpClass(cls): def setUpClass(cls):
@@ -458,6 +465,39 @@ class TestEntryPointDispatch(CustomTestCase):
rtol=5e-2, rtol=5e-2,
) )
def test_flashinfer_backend_matches_triton(self):
"""FlashInfer SM120 sparse MLA decode matches Triton reference."""
import importlib
if importlib.util.find_spec("flashinfer.sparse_mla_sm120") is None:
self.skipTest("FlashInfer SM120 sparse MLA not available")
k_cache, _ = _build_kvcache(4, 64, device=self.device, seed=5)
q, indices = _build_q_indices(1, 4, 32, 4, 64, device=self.device, seed=13)
topk_length = torch.tensor([32], dtype=torch.int32, device=self.device)
kwargs = dict(
q=q,
k_cache=k_cache,
indices=indices,
topk_length=topk_length,
attn_sink=None,
head_dim_v=_D,
softmax_scale=_D**-0.5,
)
with mock.patch.object(fmod, "_sm120_default_backend", "triton"):
out_triton, _ = flash_mla_with_kvcache_sm120(**kwargs)
with mock.patch.object(fmod, "_sm120_default_backend", "flashinfer"):
out_fi, _ = flash_mla_with_kvcache_sm120(**kwargs)
torch.testing.assert_close(
out_fi.to(torch.float32),
out_triton.to(torch.float32),
atol=5e-2,
rtol=5e-2,
)
if __name__ == "__main__": if __name__ == "__main__":
import sys import sys