Files
wehub-resource-sync 94057c3d3e
PR Test (NPU) / check-changes (push) Has been cancelled
PR Test (NPU) / pr-gate (push) Has been cancelled
PR Test (NPU) / set-image-config (push) Has been cancelled
PR Test (NPU) / stage-b-test-1-npu-a2 (0) (push) Has been cancelled
PR Test (NPU) / stage-b-test-1-npu-a2 (1) (push) Has been cancelled
PR Test (NPU) / stage-b-test-2-npu-a2 (0) (push) Has been cancelled
PR Test (NPU) / stage-b-test-2-npu-a2 (1) (push) Has been cancelled
PR Test (NPU) / stage-b-test-4-npu-a3 (push) Has been cancelled
PR Test (NPU) / stage-b-test-16-npu-a3 (push) Has been cancelled
PR Test (NPU) / multimodal-gen-test-1-npu-a3 (push) Has been cancelled
PR Test (NPU) / multimodal-gen-test-2-npu-a3 (push) Has been cancelled
PR Test (Arm64) / pr-gate (push) Has been cancelled
PR Test (Arm64) / check-changes (push) Has been cancelled
PR Test (Arm64) / build-test (push) Has been cancelled
PR Test (sgl-router) / gate (push) Has been cancelled
PR Test (sgl-router) / tier-1 — lint (push) Has been cancelled
PR Test (sgl-router) / tier-2 — build + test (push) Has been cancelled
PR Test (sgl-router) / tier-3 — docker (placeholder) (push) Has been cancelled
PR Test (sgl-router) / tier-3 — k8s integration (push) Has been cancelled
PR Test (sgl-router) / tier-3 — e2e (push) Has been cancelled
PR Test (sgl-router) / finish (push) Has been cancelled
PR Test (NPU) / single-node-poc (map[name:qwen3_6_27b_w8a8_1p_in64k_out1k_50ms runner:linux-aarch64-a3-2 test_case:test/registered/ascend/performance/qwen3_6_27b/test_npu_qwen3_6_27b_w8a8_1p_in64k_out1k_50ms.py test_type:perf]) (push) Has been cancelled
PR Test (NPU) / pr-test-npu-finish (push) Has been cancelled
PR Test (Xeon) / pr-gate (push) Has been cancelled
PR Test (Xeon) / check-changes (push) Has been cancelled
PR Test (Xeon) / build-test (, xeon-gnr, base-b-test-cpu) (push) Has been cancelled
PR Test (XPU) / check-changes (push) Has been cancelled
PR Test (XPU) / pr-gate (push) Has been cancelled
PR Test (XPU) / stage-a-test-1-gpu-xpu (push) Has been cancelled
PR Test (XPU) / wait-for-stage-a (push) Has been cancelled
PR Test (XPU) / stage-b-test-1-gpu-xpu (push) Has been cancelled
PR Test (XPU) / finish (push) Has been cancelled
CI Model Inventory / build-inventory (push) Has been cancelled
Lint / lint (push) Has been cancelled
PR Benchmark (SMG Components) / Benchmark Compilation Check (push) Has been cancelled
PR Benchmark (SMG Components) / Benchmark - Manual Policy (push) Has been cancelled
PR Benchmark (SMG Components) / Benchmark - Request Processing (push) Has been cancelled
PR Benchmark (SMG Components) / Benchmark Summary (push) Has been cancelled
PR Test (SMG) / build-wheel (push) Has been cancelled
Release SGLang Model Gateway to PyPI / build on windows (x86_64 - auto) (push) Has been cancelled
Release SGLang Model Gateway to PyPI / build on macos (x86_64 - auto) (push) Has been cancelled
PR Test (SMG) / python-unit-tests (push) Has been cancelled
PR Test (SMG) / unit-tests (push) Has been cancelled
PR Test (SMG) / benchmarks (push) Has been cancelled
PR Test (SMG) / chat-completions (push) Has been cancelled
PR Test (SMG) / chat-completions-4gpu (push) Has been cancelled
PR Test (SMG) / e2e (push) Has been cancelled
PR Test (SMG) / docker-build-test (push) Has been cancelled
PR Test (SMG) / k8s-integration (push) Has been cancelled
PR Test (SMG) / finish (push) Has been cancelled
PR Test (SMG) / summarize-benchmarks (push) Has been cancelled
Release SGLang Model Gateway Docker Image / publish (push) Has been cancelled
Release SGLang Model Gateway to PyPI / build on macos (aarch64 - auto) (push) Has been cancelled
Release SGLang Model Gateway to PyPI / build on linux (aarch64 - auto) (push) Has been cancelled
Release SGLang Model Gateway to PyPI / build on linux (x86_64 - auto) (push) Has been cancelled
Release SGLang Model Gateway to PyPI / build on linux (aarch64 - musllinux_1_1) (push) Has been cancelled
Release SGLang Model Gateway to PyPI / build on linux (x86_64 - musllinux_1_1) (push) Has been cancelled
Release SGLang Model Gateway to PyPI / Build SDist (push) Has been cancelled
Release SGLang Model Gateway to PyPI / Upload to PyPI (push) Has been cancelled
Release SGLang Kernels / build-cu129-matrix (aarch64, 12.9, 3.10, arm-kernel-build-node) (push) Has been cancelled
Release SGLang Kernels / build-cu129-matrix (x86_64, 12.9, 3.10, x64-kernel-build-node) (push) Has been cancelled
Release SGLang Kernels / release-cu129 (push) Has been cancelled
Release SGLang Kernels / build-cu130-matrix (aarch64, 13.0, 3.10, arm-kernel-build-node) (push) Has been cancelled
Release SGLang Kernels / build-cu130-matrix (x86_64, 13.0, 3.10, x64-kernel-build-node) (push) Has been cancelled
Release SGLang Kernels / release-cu130 (push) Has been cancelled
Release SGLang Kernels / build-rocm-matrix (3.10, 700) (push) Has been cancelled
Release SGLang Kernels / build-rocm-matrix (3.10, 720) (push) Has been cancelled
Release SGLang Kernels / release-rocm700 (push) Has been cancelled
Release SGLang Kernels / release-rocm720 (push) Has been cancelled
Release SGLang Kernels / build-musa43 (43, 3.10) (push) Has been cancelled
Release SGLang Kernels / release-musa43 (push) Has been cancelled
chore: import upstream snapshot with attribution
2026-07-13 12:38:16 +08:00

161 lines
5.1 KiB
Python

"""Triton sparse-MLA forward for the DSA fp8 prefill path.
A per-query flash-attention kernel over the indexer-selected topk KV. On
gfx950 this is ~1.6x faster than the TileLang partial+combine kernel for the
prefill regime (n_groups=1): the attention tile is tiny (M=16 heads = one
16x16 MFMA), so a small-warp per-program kernel avoids the intra-block
coordination overhead of the 256-thread TileLang block.
"""
import torch
import triton
import triton.language as tl
from sglang.srt.layers.quantization.fp8_kernel import is_fp8_fnuz
_IS_FNUZ = is_fp8_fnuz()
_FP8_MAX = 240.0 if _IS_FNUZ else 448.0
def _prune_configs(configs, named_args, **kwargs):
"""Drop configs whose KV tile exceeds topk (pure waste)."""
topk = named_args["topk"]
keep = [c for c in configs if c.kwargs["BLOCK_N"] <= topk]
return keep or [configs[0]]
# The best (BLOCK_N, num_warps, num_stages) is shape- and arch-sensitive, so
# autotune over a grid keyed on the attention shape. Benchmarked once per key
# (a one-time stall on the first prefill of each new shape), then cached.
_AUTOTUNE_CONFIGS = [
triton.Config({"BLOCK_N": bn}, num_warps=w, num_stages=ns)
for bn in (32, 64, 128)
for w in (1, 2, 4)
for ns in (1, 2)
]
@triton.autotune(
configs=_AUTOTUNE_CONFIGS,
key=["topk", "H", "DIM"],
prune_configs_by={"early_config_prune": _prune_configs},
)
@triton.jit
def _sparse_mla_fwd_kernel(
q_nope_ptr,
q_rope_ptr,
kv_ptr,
idx_ptr,
o_ptr,
sm_scale,
fp8_max,
topk,
H: tl.constexpr,
DIM: tl.constexpr,
D_V: tl.constexpr,
D_TAIL: tl.constexpr,
BLOCK_N: tl.constexpr,
):
s_i = tl.program_id(0)
h = tl.arange(0, H)
dv = tl.arange(0, D_V)
dt = tl.arange(0, D_TAIL)
# q is read as two separate tensors (q_nope width D_V, q_rope width D_TAIL):
# the upstream concat into a single [.., DIM] tensor is skipped since this
# kernel splits q into main/tail anyway.
q_main = tl.load(q_nope_ptr + s_i * H * D_V + h[:, None] * D_V + dv[None, :]).to(
q_nope_ptr.dtype.element_ty
) # [H, D_V]
q_tail = tl.load(
q_rope_ptr + s_i * H * D_TAIL + h[:, None] * D_TAIL + dt[None, :]
).to(
q_nope_ptr.dtype.element_ty
) # [H, D_TAIL]
m_i = tl.full([H], -float("inf"), tl.float32)
l_i = tl.zeros([H], tl.float32)
acc = tl.zeros([H, D_V], tl.float32)
n = tl.arange(0, BLOCK_N)
for k0 in range(0, topk, BLOCK_N):
kmask = (k0 + n) < topk
idx = tl.load(idx_ptr + s_i * topk + k0 + n, mask=kmask, other=-1)
valid = (idx >= 0) & kmask
page = tl.where(valid, idx, 0)
kbase = kv_ptr + page[:, None] * DIM
kv_main = tl.load(kbase + dv[None, :], mask=valid[:, None], other=0.0).to(
q_nope_ptr.dtype.element_ty
) # [BLOCK_N, D_V] -- reused as V
kv_tail = tl.load(
kbase + (D_V + dt)[None, :], mask=valid[:, None], other=0.0
).to(
q_nope_ptr.dtype.element_ty
) # [BLOCK_N, D_TAIL]
qk = tl.dot(q_main, tl.trans(kv_main)).to(tl.float32)
qk += tl.dot(q_tail, tl.trans(kv_tail)).to(tl.float32)
qk = qk * sm_scale
qk = tl.where(valid[None, :], qk, -float("inf"))
m_new = tl.maximum(m_i, tl.max(qk, axis=1))
# Guard an all-masked row (m_new == -inf): shift by 0 instead so that
# exp(-inf - 0) = 0 rather than exp(-inf + inf) = NaN. Identical to
# m_new whenever the row has >=1 valid key.
m_safe = tl.where(m_new == -float("inf"), 0.0, m_new)
alpha = tl.exp(m_i - m_safe)
p = tl.exp(qk - m_safe[:, None]) # [H, BLOCK_N]
l_i = l_i * alpha + tl.sum(p, axis=1)
p_fp8 = (p * fp8_max).to(q_nope_ptr.dtype.element_ty)
pv = tl.dot(p_fp8, kv_main).to(tl.float32) * (1.0 / fp8_max)
acc = acc * alpha[:, None] + pv
m_i = m_new
l_safe = tl.where(l_i == 0.0, 1.0, l_i)
acc = acc / l_safe[:, None]
tl.store(
o_ptr + s_i * H * D_V + h[:, None] * D_V + dv[None, :],
acc.to(o_ptr.dtype.element_ty),
)
def triton_sparse_mla_fwd(
q_nope: torch.Tensor,
q_rope: torch.Tensor,
kv: torch.Tensor,
indices: torch.Tensor,
sm_scale: float,
d_v: int = 512,
) -> torch.Tensor:
"""q_nope: [seq, H, d_v] fp8, q_rope: [seq, H, dim-d_v] fp8,
kv: [num_pages, 1, dim] fp8, indices: [seq, 1, topk].
Reads q from the two un-concatenated tensors directly (no q_nope/q_rope
concat). Returns [1, seq, H, d_v] bf16 to match tilelang_sparse_fwd.
"""
seq, H, d_v_in = q_nope.shape
assert d_v_in == d_v
d_tail = q_rope.shape[-1]
dim = kv.shape[-1]
topk = indices.shape[-1]
q_nope = q_nope.contiguous()
q_rope = q_rope.contiguous()
out = torch.empty(seq, H, d_v, device=q_nope.device, dtype=torch.bfloat16)
# BLOCK_N / num_warps / num_stages are chosen by @triton.autotune.
_sparse_mla_fwd_kernel[(seq,)](
q_nope,
q_rope,
kv,
indices,
out,
sm_scale,
_FP8_MAX,
topk,
H=H,
DIM=dim,
D_V=d_v,
D_TAIL=d_tail,
)
return out.unsqueeze(0)