[SM120] Add FlashInfer sparse MLA decode for DSv4-Flash (#27455)
Co-authored-by: AliceChenyy <alicechenyy@users.noreply.github.com>
This commit is contained in:
@@ -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)
|
||||||
|
|||||||
+40
@@ -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
|
||||||
Reference in New Issue
Block a user