[DSA] Q8KV8 FP8 Sparse Prefill on GLM-5.2 & DeepSeek-V3.2: Q8-Path & Shared-Path Optimizations (#31888)
Co-authored-by: Claude Fable 5 <noreply@anthropic.com>
This commit is contained in:
co-authored by
Claude Fable 5
parent
4f51dad1da
commit
e4a40a71f8
+262
@@ -0,0 +1,262 @@
|
||||
#!/usr/bin/env python3
|
||||
"""Microbenchmark: Q8KV8 sparse-prefill KV gather overhaul.
|
||||
|
||||
Compares the legacy gather (``gather_dequant_requant_fp8_paged_legacy``:
|
||||
fresh ``torch.zeros`` destination + one program per (token, 128-elem
|
||||
slice)) against the new gather (``gather_dequant_requant_fp8_paged``:
|
||||
no pre-zeroing needed, fused pad-row zero-fill, TOKENS_PER_PROG tokens
|
||||
per program with 16B-vectorized access), in three flavors:
|
||||
|
||||
legacy torch.zeros alloc + legacy kernel (baseline)
|
||||
new_alloc torch.empty alloc + vectorized kernel (vectorized copy only)
|
||||
new_cached persistent grow-only buffer + vec kernel (changes 1 + 2,
|
||||
= production path)
|
||||
|
||||
For each scenario it checks BIT-EXACT equality of the fp8 output bytes
|
||||
(``torch.equal`` on ``uint8`` views; the requant scale is the identity
|
||||
scalar 1.0 on this path, so the buffer is the entire output), including
|
||||
the `topk` zero landing-pad rows, then reports us/call and effective
|
||||
TB/s.
|
||||
|
||||
Shapes model GLM / DeepSeek-V3.2 DSA prefill on one rank: d = 576 fp8
|
||||
out (512 nope + 64 rope), 656 B/token paged cache rows (512 nope fp8 +
|
||||
16 B f32 group scales + 128 B bf16 rope), page_size 64, topk 2048.
|
||||
NOTE: gather traffic scales with kv_len (= number of gathered KV rows
|
||||
= len(page_table_1_flattened)), NOT with s_q; s_q below only labels the
|
||||
chunk that a scenario represents. The GLM-5.2 il=64k profile point
|
||||
(112.0 us/call, ~0.7 TB/s effective) corresponds to the
|
||||
(s_q=4096, kv_len=65536) row.
|
||||
|
||||
Usage (single GPU, < 2 min):
|
||||
python benchmark/kernels/deepseek/benchmark_q8kv8_kv_gather.py [--device cuda:0]
|
||||
[--iters 200] [--warmup 20]
|
||||
|
||||
Also runs correctness-only edge cases: ragged tail (kv_len % 4 != 0),
|
||||
extra_rows=0, num_tokens=0, and a cached-buffer SHRINK reuse (big call
|
||||
then small call) that proves stale bytes from the earlier, larger call
|
||||
cannot leak into the smaller call's pad rows.
|
||||
"""
|
||||
|
||||
import argparse
|
||||
import importlib.util
|
||||
import sys
|
||||
from pathlib import Path
|
||||
|
||||
import torch
|
||||
|
||||
_REPO_ROOT = Path(__file__).resolve().parents[3]
|
||||
_MOD_PATH = _REPO_ROOT / "python/sglang/kernels/ops/attention/dsa/dequant_k_cache.py"
|
||||
|
||||
# Import the module straight from its file so the benchmark stays
|
||||
# standalone (no sglang package import side effects; needs only
|
||||
# torch + triton).
|
||||
_spec = importlib.util.spec_from_file_location("dequant_k_cache", _MOD_PATH)
|
||||
_mod = importlib.util.module_from_spec(_spec)
|
||||
_spec.loader.exec_module(_mod)
|
||||
|
||||
gather_new = _mod.gather_dequant_requant_fp8_paged
|
||||
gather_legacy = _mod.gather_dequant_requant_fp8_paged_legacy
|
||||
|
||||
PAGE_SIZE = 64
|
||||
DIM_QUANT = 656 # 512 nope fp8 + 16 scale bytes + 128 rope bytes
|
||||
OUT_DIM = 576 # 512 nope + 64 rope, fp8
|
||||
TOPK = 2048
|
||||
|
||||
# (s_q label, kv_len = gathered rows). kv_len drives the bytes moved.
|
||||
SCENARIOS = [
|
||||
(512, 8192),
|
||||
(2048, 32768),
|
||||
(4096, 65536),
|
||||
]
|
||||
|
||||
|
||||
def build_paged_kv_pool(pool_tokens: int, device: str) -> torch.Tensor:
|
||||
"""Synthetic paged fp8 KV cache: [pool_tokens, 1, 656] fp8_e4m3fn."""
|
||||
g = torch.Generator(device=device).manual_seed(0)
|
||||
nope = (torch.randn(pool_tokens, 512, generator=g, device=device) * 2.0).to(
|
||||
torch.float8_e4m3fn
|
||||
)
|
||||
# Positive, realistically small per-group dequant scales.
|
||||
scales = (torch.rand(pool_tokens, 4, generator=g, device=device) * 0.05 + 1e-3).to(
|
||||
torch.float32
|
||||
)
|
||||
rope = torch.randn(pool_tokens, 64, generator=g, device=device).to(torch.bfloat16)
|
||||
|
||||
raw = torch.empty(pool_tokens, DIM_QUANT, dtype=torch.uint8, device=device)
|
||||
raw[:, :512] = nope.view(torch.uint8)
|
||||
raw[:, 512:528] = scales.view(torch.uint8)
|
||||
raw[:, 528:] = rope.view(torch.uint8)
|
||||
return raw.view(torch.float8_e4m3fn).view(pool_tokens, 1, DIM_QUANT)
|
||||
|
||||
|
||||
def build_page_table_flattened(
|
||||
kv_lens, pool_tokens: int, device: str, seed: int = 1
|
||||
) -> torch.Tensor:
|
||||
"""Realistic page_table_1_flattened: per request, random distinct
|
||||
64-token pages, tokens contiguous within a page (production paged
|
||||
layout), requests concatenated."""
|
||||
n_pool_pages = pool_tokens // PAGE_SIZE
|
||||
g = torch.Generator(device="cpu").manual_seed(seed)
|
||||
parts = []
|
||||
for kv_len in kv_lens:
|
||||
n_pages = (kv_len + PAGE_SIZE - 1) // PAGE_SIZE
|
||||
assert n_pages <= n_pool_pages, "pool too small for scenario"
|
||||
pages = torch.randperm(n_pool_pages, generator=g)[:n_pages]
|
||||
toks = (pages[:, None] * PAGE_SIZE + torch.arange(PAGE_SIZE)[None, :]).reshape(
|
||||
-1
|
||||
)[:kv_len]
|
||||
parts.append(toks)
|
||||
return torch.cat(parts).to(torch.int32).to(device)
|
||||
|
||||
|
||||
class CachedGather:
|
||||
"""Mimics the dsa_backend production path: persistent grow-only fp8
|
||||
destination buffer, gather zero-fills only the pad tail in-kernel."""
|
||||
|
||||
def __init__(self):
|
||||
self.buf = None
|
||||
|
||||
def __call__(self, pool, pt, extra_rows):
|
||||
total = pt.shape[0] + extra_rows
|
||||
if self.buf is None or self.buf.shape[0] < total:
|
||||
self.buf = torch.empty(
|
||||
(total, OUT_DIM), dtype=torch.float8_e4m3fn, device=pool.device
|
||||
)
|
||||
return gather_new(pool, pt, extra_rows=extra_rows, out=self.buf[:total])
|
||||
|
||||
|
||||
def assert_bit_exact(ref: torch.Tensor, got: torch.Tensor, what: str):
|
||||
assert ref.shape == got.shape, f"{what}: shape {got.shape} != {ref.shape}"
|
||||
ok = torch.equal(
|
||||
ref.contiguous().view(torch.uint8), got.contiguous().view(torch.uint8)
|
||||
)
|
||||
assert ok, f"{what}: fp8 bytes NOT bit-exact"
|
||||
|
||||
|
||||
def bench_us(fn, warmup: int, iters: int) -> float:
|
||||
for _ in range(warmup):
|
||||
fn()
|
||||
torch.cuda.synchronize()
|
||||
start = torch.cuda.Event(enable_timing=True)
|
||||
end = torch.cuda.Event(enable_timing=True)
|
||||
start.record()
|
||||
for _ in range(iters):
|
||||
fn()
|
||||
end.record()
|
||||
torch.cuda.synchronize()
|
||||
return start.elapsed_time(end) * 1000.0 / iters # ms -> us
|
||||
|
||||
|
||||
def main():
|
||||
ap = argparse.ArgumentParser()
|
||||
ap.add_argument("--device", default="cuda:0")
|
||||
ap.add_argument("--iters", type=int, default=200)
|
||||
ap.add_argument("--warmup", type=int, default=20)
|
||||
args = ap.parse_args()
|
||||
|
||||
if not torch.cuda.is_available():
|
||||
print("CUDA not available; this microbench needs 1 GPU.", file=sys.stderr)
|
||||
sys.exit(1)
|
||||
torch.cuda.set_device(args.device)
|
||||
dev = args.device
|
||||
|
||||
pool_tokens = 131072 # 2x the largest kv_len; 131072*656 B ~ 86 MB
|
||||
pool = build_paged_kv_pool(pool_tokens, dev)
|
||||
|
||||
print(f"device={dev} ({torch.cuda.get_device_name(dev)})")
|
||||
print(f"pool: {pool_tokens} tokens x {DIM_QUANT} B, topk={TOPK}")
|
||||
print()
|
||||
|
||||
# ------------------------------------------------------------------
|
||||
# Correctness edge cases (not timed)
|
||||
# ------------------------------------------------------------------
|
||||
print("== correctness edge cases ==")
|
||||
cached_edge = CachedGather()
|
||||
for kv_len, extra in [(1234, TOPK), (8192, 0), (0, TOPK), (63, 17)]:
|
||||
pt = build_page_table_flattened([kv_len], pool_tokens, dev, seed=7)
|
||||
ref = gather_legacy(pool, pt, extra_rows=extra)
|
||||
got_alloc = gather_new(pool, pt, extra_rows=extra)
|
||||
got_cached = cached_edge(pool, pt, extra_rows=extra)
|
||||
assert_bit_exact(ref, got_alloc, f"kv_len={kv_len},extra={extra} new_alloc")
|
||||
assert_bit_exact(ref, got_cached, f"kv_len={kv_len},extra={extra} new_cached")
|
||||
# Pad rows must be exactly zero bytes.
|
||||
if extra > 0:
|
||||
pad = got_cached[kv_len:].view(torch.uint8)
|
||||
assert int(pad.max()) == 0 if pad.numel() else True
|
||||
print(f" kv_len={kv_len:6d} extra_rows={extra:5d}: bit-exact OK")
|
||||
|
||||
# Cached-buffer SHRINK reuse: big call dirties the buffer, then a
|
||||
# smaller call must still produce zero pad rows (stale-data test for
|
||||
# the grow-only buffer + tail-only zeroing invariant).
|
||||
cached_shrink = CachedGather()
|
||||
pt_big = build_page_table_flattened([65536], pool_tokens, dev, seed=11)
|
||||
cached_shrink(pool, pt_big, TOPK)
|
||||
pt_small = build_page_table_flattened([4096], pool_tokens, dev, seed=13)
|
||||
ref_small = gather_legacy(pool, pt_small, extra_rows=TOPK)
|
||||
got_small = cached_shrink(pool, pt_small, TOPK)
|
||||
assert_bit_exact(ref_small, got_small, "shrink-reuse (65536 -> 4096)")
|
||||
print(" shrink-reuse 65536 -> 4096 rows: pad rows clean, bit-exact OK")
|
||||
print()
|
||||
|
||||
# ------------------------------------------------------------------
|
||||
# Timed scenarios
|
||||
# ------------------------------------------------------------------
|
||||
print("== timing ==")
|
||||
print(
|
||||
"metrics: us/call = mean wall time per gather call incl. any alloc/"
|
||||
"zero-fill (LOWER = faster); eff TB/s = payload (656+4 B/token read"
|
||||
" + 576 B/row written incl. pad) / time (HIGHER = faster);"
|
||||
" speedup = legacy_us / variant_us (>1 = faster than legacy)."
|
||||
)
|
||||
header = (
|
||||
f"{'s_q':>5} {'kv_len':>7} {'variant':>10} {'us/call':>9} "
|
||||
f"{'eff TB/s':>9} {'speedup':>8}"
|
||||
)
|
||||
print(header)
|
||||
print("-" * len(header))
|
||||
|
||||
for s_q, kv_len in SCENARIOS:
|
||||
pt = build_page_table_flattened([kv_len], pool_tokens, dev, seed=s_q)
|
||||
total_rows = kv_len + TOPK
|
||||
payload_bytes = kv_len * (DIM_QUANT + 4) + total_rows * OUT_DIM
|
||||
|
||||
cached = CachedGather()
|
||||
variants = [
|
||||
("legacy", lambda: gather_legacy(pool, pt, extra_rows=TOPK)),
|
||||
("new_alloc", lambda: gather_new(pool, pt, extra_rows=TOPK)),
|
||||
("new_cached", lambda: cached(pool, pt, TOPK)),
|
||||
]
|
||||
|
||||
# Bit-exactness at the benchmarked shape before timing.
|
||||
ref = gather_legacy(pool, pt, extra_rows=TOPK)
|
||||
for name, fn in variants[1:]:
|
||||
assert_bit_exact(ref, fn(), f"s_q={s_q} {name}")
|
||||
del ref
|
||||
|
||||
legacy_us = None
|
||||
for name, fn in variants:
|
||||
us = bench_us(fn, args.warmup, args.iters)
|
||||
tbps = payload_bytes / (us * 1e-6) / 1e12
|
||||
if name == "legacy":
|
||||
legacy_us = us
|
||||
speedup = "1.00x"
|
||||
else:
|
||||
speedup = f"{legacy_us / us:.2f}x"
|
||||
print(
|
||||
f"{s_q:>5} {kv_len:>7} {name:>10} {us:>9.1f} "
|
||||
f"{tbps:>9.3f} {speedup:>8}"
|
||||
)
|
||||
print()
|
||||
|
||||
print(
|
||||
"note: legacy additionally writes a full-buffer zero fill "
|
||||
f"({total_rows * OUT_DIM / 1e6:.1f} MB at the largest shape) that is "
|
||||
"NOT counted in its payload bytes; its true HW bandwidth is higher "
|
||||
"than the eff TB/s shown, which is exactly why us/call is the "
|
||||
"decision metric."
|
||||
)
|
||||
print("ALL CHECKS PASSED")
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
+486
@@ -0,0 +1,486 @@
|
||||
#!/usr/bin/env python3
|
||||
"""Microbenchmark: Q8KV8 sparse-prefill q-prep — old path vs born-fp8 fused path.
|
||||
|
||||
Old path (production default):
|
||||
1. q_nope_out = torch.bmm(q_nope.transpose(0, 1), w_kc).transpose(0, 1)
|
||||
(cublas bf16 bmm, writes bf16 [H, T, N] to DRAM)
|
||||
2. concat_and_cast_q_fp8_pad(q_fp8, q_nope_out, q_rope, H)
|
||||
(Triton: re-reads the bf16 bmm output + q_rope, writes fp8 [T, H, N+R])
|
||||
|
||||
New path (SGLANG_ENABLE_DSA_Q8KV8_BORN_FP8_Q):
|
||||
absorbed_bmm_concat_cast_q_fp8(q_fp8, q_nope, w_kc, q_rope, H)
|
||||
(one Triton kernel: bmm + concat + fp8 cast; the bf16 q_nope_out and the
|
||||
standalone concat-cast launch disappear)
|
||||
|
||||
Non-power-of-2 K (GLM 192) K-dimension codegen variants (A/B'd here; all keep
|
||||
the identical fp32 -> bf16 -> fp8 epilogue, see cache_ops.py):
|
||||
loop : split-K loop, BLOCK_K=64 x 3 (the original K=192 path)
|
||||
two_dot : preload a once as 128+64 tiles, two chained tl.dot, no K-loop
|
||||
three_dot : preload a once as 3 x 64 tiles, three chained tl.dot
|
||||
(same fp32 add order as `loop`, loads hoisted)
|
||||
pad : single tl.dot at BLOCK_K=256 with zero-masked k tail
|
||||
single_k : single tl.dot at BLOCK_K=192 -- documents the Triton
|
||||
non-power-of-2 tl.arange limitation (compile fails <= 3.5.x)
|
||||
Power-of-2 K (DeepSeek 128) collapses every variant to the same single-dot
|
||||
fast path, so only one NEW row is shown there.
|
||||
|
||||
Shapes (both models: N = kv_lora_rank = 512, R = qk_rope_head_dim = 64;
|
||||
K = qk_nope_head_dim differs per model):
|
||||
GLM-5.2: heads = 64, K = 192 (w_kc [64, 192, 512]; DP attention,
|
||||
per-rank full heads; K=192 exercises the kernel's split-K path)
|
||||
DS-V3.2: heads = 128, K = 128 (power-of-2 K, preload-once fast path)
|
||||
|
||||
Metric conventions:
|
||||
* time is reported in microseconds per call (us/call) — LOWER = FASTER.
|
||||
* bandwidth is analytic-bytes / time in GB/s — HIGHER = BETTER.
|
||||
* "speedup x" = old_time / new_time — >1.0 means the NEW path is faster.
|
||||
|
||||
Correctness:
|
||||
* rope half must be BIT-EXACT (same bf16 source, same Triton conversion).
|
||||
* nope half: same rounding stages (fp32 accum -> bf16 -> fp8) but a
|
||||
different GEMM accumulation order than cublas -> near- but not
|
||||
guaranteed bit-exact. We report the bitwise-match fraction, the max
|
||||
dequantized |diff|, and which path lands closer to an fp64 reference.
|
||||
|
||||
Usage (single GPU):
|
||||
python benchmark/kernels/deepseek/benchmark_q8kv8_q_prep.py # both model shapes
|
||||
python benchmark/kernels/deepseek/benchmark_q8kv8_q_prep.py --tokens 8192 --iters 300
|
||||
python benchmark/kernels/deepseek/benchmark_q8kv8_q_prep.py --variants two_dot,pad
|
||||
python benchmark/kernels/deepseek/benchmark_q8kv8_q_prep.py --sweep # + tile/warp sweep
|
||||
python benchmark/kernels/deepseek/benchmark_q8kv8_q_prep.py --rounding-study
|
||||
"""
|
||||
|
||||
import argparse
|
||||
|
||||
import torch
|
||||
|
||||
from sglang.kernels.ops.kvcache.cache_ops import (
|
||||
absorbed_bmm_concat_cast_q_fp8,
|
||||
concat_and_cast_q_fp8_pad,
|
||||
)
|
||||
|
||||
N_LORA = 512 # kv_lora_rank (post-absorb q_nope dim; "d_nope" at the kernel)
|
||||
R_ROPE = 64 # qk_rope_head_dim
|
||||
|
||||
|
||||
def make_inputs(
|
||||
num_tokens: int,
|
||||
num_heads: int,
|
||||
k_nope: int,
|
||||
device,
|
||||
seed: int,
|
||||
magnitude: float,
|
||||
):
|
||||
g = torch.Generator(device=device).manual_seed(seed)
|
||||
# Production layout: q = q_b_proj output [T, H, K+R] bf16; q_nope/q_rope are
|
||||
# strided views of it (rope applied in-place on the q_rope slice).
|
||||
q = (
|
||||
torch.randn(
|
||||
(num_tokens, num_heads, k_nope + R_ROPE),
|
||||
generator=g,
|
||||
device=device,
|
||||
dtype=torch.float32,
|
||||
)
|
||||
* magnitude
|
||||
).to(torch.bfloat16)
|
||||
q_nope = q[..., :k_nope]
|
||||
q_rope = q[..., k_nope:]
|
||||
# Production w_kc layout: [H, K, N] with strides (K*N, 1, K) (N-major), the
|
||||
# result of w_kc.transpose(1, 2).contiguous().transpose(1, 2) at load.
|
||||
w_base = (
|
||||
torch.randn(
|
||||
(num_heads, N_LORA, k_nope),
|
||||
generator=g,
|
||||
device=device,
|
||||
dtype=torch.float32,
|
||||
)
|
||||
/ (k_nope**0.5)
|
||||
).to(torch.bfloat16)
|
||||
w_kc = w_base.transpose(1, 2)
|
||||
return q, q_nope, q_rope, w_kc
|
||||
|
||||
|
||||
def old_path(q_fp8, q_nope, w_kc, q_rope, num_heads):
|
||||
q_nope_out = torch.bmm(q_nope.transpose(0, 1), w_kc).transpose(0, 1)
|
||||
concat_and_cast_q_fp8_pad(q_fp8, q_nope_out, q_rope, num_heads)
|
||||
|
||||
|
||||
def old_path_bmm_only(q_nope, w_kc):
|
||||
return torch.bmm(q_nope.transpose(0, 1), w_kc)
|
||||
|
||||
|
||||
def new_path(q_fp8, q_nope, w_kc, q_rope, num_heads, **kw):
|
||||
absorbed_bmm_concat_cast_q_fp8(q_fp8, q_nope, w_kc, q_rope, num_heads, **kw)
|
||||
|
||||
|
||||
# Non-power-of-2-K variants, in bench order (power-of-2 K collapses to "auto").
|
||||
ALL_VARIANTS = ["loop", "two_dot", "three_dot", "pad", "single_k"]
|
||||
|
||||
# (block_m, block_n, num_warps, num_stages) sweep grid; num_stages 0 = Triton
|
||||
# default. N=512 is a multiple of every block_n here; block_m stays power of 2.
|
||||
SWEEP_TILES = [
|
||||
(64, 128, 4, 0), # kernel default
|
||||
(64, 128, 8, 0),
|
||||
(64, 128, 4, 2),
|
||||
(64, 128, 4, 4),
|
||||
(128, 128, 4, 0),
|
||||
(128, 128, 8, 0),
|
||||
(64, 256, 8, 0),
|
||||
(128, 256, 8, 0),
|
||||
(32, 128, 4, 0),
|
||||
(64, 64, 4, 0),
|
||||
(128, 64, 8, 0),
|
||||
]
|
||||
|
||||
|
||||
def time_fn(fn, iters: int, warmup: int) -> float:
|
||||
"""Median wall time of fn() in microseconds per call (lower = faster)."""
|
||||
for _ in range(warmup):
|
||||
fn()
|
||||
torch.cuda.synchronize()
|
||||
times = []
|
||||
start = torch.cuda.Event(enable_timing=True)
|
||||
end = torch.cuda.Event(enable_timing=True)
|
||||
chunk = 10
|
||||
for _ in range(max(1, iters // chunk)):
|
||||
start.record()
|
||||
for _ in range(chunk):
|
||||
fn()
|
||||
end.record()
|
||||
torch.cuda.synchronize()
|
||||
times.append(start.elapsed_time(end) * 1e3 / chunk) # ms -> us
|
||||
times.sort()
|
||||
return times[len(times) // 2]
|
||||
|
||||
|
||||
def analytic_bytes(num_tokens: int, num_heads: int, k_nope: int):
|
||||
"""(old_bytes, new_bytes) of DRAM traffic per call, analytic lower bound."""
|
||||
t, h = num_tokens, num_heads
|
||||
a_read = t * h * k_nope * 2 # q_nope bf16
|
||||
w_read = h * k_nope * N_LORA * 2 # w_kc bf16
|
||||
nope_bf16 = t * h * N_LORA * 2 # bmm bf16 out (written then re-read)
|
||||
rope_read = t * h * R_ROPE * 2 # q_rope bf16
|
||||
fp8_write = t * h * (N_LORA + R_ROPE) # q_fp8 out
|
||||
old = (a_read + w_read + nope_bf16) + (nope_bf16 + rope_read + fp8_write)
|
||||
new = a_read + w_read + rope_read + fp8_write
|
||||
return old, new
|
||||
|
||||
|
||||
def make_check_ctx(num_tokens, num_heads, k_nope, device, seed, magnitude):
|
||||
"""Fresh inputs + old-path fp8 output + fp64 bmm reference (once/config)."""
|
||||
q, q_nope, q_rope, w_kc = make_inputs(
|
||||
num_tokens, num_heads, k_nope, device, seed, magnitude
|
||||
)
|
||||
q_fp8_old = torch.zeros(
|
||||
(num_tokens, num_heads, N_LORA + R_ROPE),
|
||||
dtype=torch.float8_e4m3fn,
|
||||
device=device,
|
||||
)
|
||||
old_path(q_fp8_old, q_nope, w_kc, q_rope, num_heads)
|
||||
torch.cuda.synchronize()
|
||||
# fp64 reference: which path's fp8 lands closer to the exact bmm?
|
||||
ref = torch.bmm(
|
||||
q_nope.transpose(0, 1).to(torch.float64), w_kc.to(torch.float64)
|
||||
).transpose(0, 1)
|
||||
err_old = (q_fp8_old[..., :N_LORA].to(torch.float64) - ref).abs()
|
||||
return {
|
||||
"q_nope": q_nope,
|
||||
"q_rope": q_rope,
|
||||
"w_kc": w_kc,
|
||||
"num_heads": num_heads,
|
||||
"q_fp8_old": q_fp8_old,
|
||||
"ref": ref,
|
||||
"meanerr_old": err_old.mean().item(),
|
||||
"maxerr_old": err_old.max().item(),
|
||||
}
|
||||
|
||||
|
||||
def check_variant(ctx, **new_kwargs):
|
||||
"""Correctness of one new-path variant vs the old path + fp64 reference."""
|
||||
q_fp8_old = ctx["q_fp8_old"]
|
||||
q_fp8_new = torch.zeros_like(q_fp8_old)
|
||||
new_path(
|
||||
q_fp8_new,
|
||||
ctx["q_nope"],
|
||||
ctx["w_kc"],
|
||||
ctx["q_rope"],
|
||||
ctx["num_heads"],
|
||||
**new_kwargs,
|
||||
)
|
||||
torch.cuda.synchronize()
|
||||
|
||||
rope_old = q_fp8_old[..., N_LORA:].view(torch.uint8)
|
||||
rope_new = q_fp8_new[..., N_LORA:].view(torch.uint8)
|
||||
rope_bitexact = bool(torch.equal(rope_old, rope_new))
|
||||
|
||||
nope_old = q_fp8_old[..., :N_LORA]
|
||||
nope_new = q_fp8_new[..., :N_LORA]
|
||||
match = (
|
||||
(nope_old.view(torch.uint8) == nope_new.view(torch.uint8)).float().mean().item()
|
||||
)
|
||||
diff = (nope_old.to(torch.float32) - nope_new.to(torch.float32)).abs()
|
||||
max_diff = diff.max().item()
|
||||
|
||||
err_new = (nope_new.to(torch.float64) - ctx["ref"]).abs()
|
||||
return {
|
||||
"rope_bitexact": rope_bitexact,
|
||||
"nope_bitwise_match_frac": match,
|
||||
"nope_max_dequant_absdiff": max_diff,
|
||||
"nope_meanerr_old_vs_fp64": ctx["meanerr_old"],
|
||||
"nope_meanerr_new_vs_fp64": err_new.mean().item(),
|
||||
"nope_maxerr_old_vs_fp64": ctx["maxerr_old"],
|
||||
"nope_maxerr_new_vs_fp64": err_new.max().item(),
|
||||
}
|
||||
|
||||
|
||||
def rounding_study(device, seed):
|
||||
"""sweep summary: fp8(bf16(x)) double round vs fp8(x) single round.
|
||||
|
||||
(Informational only — the born-fp8 kernel deliberately keeps the
|
||||
fp32->bf16->fp8 double round to match the default path's rounding stages.)
|
||||
"""
|
||||
g = torch.Generator(device=device).manual_seed(seed)
|
||||
x = torch.randn((1 << 22,), generator=g, device=device, dtype=torch.float32) * 8.0
|
||||
double = x.to(torch.bfloat16).to(torch.float8_e4m3fn)
|
||||
single = x.to(torch.float8_e4m3fn)
|
||||
mismatch = (
|
||||
(double.view(torch.uint8) != single.view(torch.uint8)).float().mean().item()
|
||||
)
|
||||
err_double = (double.to(torch.float32) - x).abs()
|
||||
err_single = (single.to(torch.float32) - x).abs()
|
||||
print("\n=== rounding study: fp32->bf16->fp8 (double) vs fp32->fp8 (single) ===")
|
||||
print(f"elements : {x.numel()}")
|
||||
print(f"byte-mismatch fraction : {mismatch:.3e} (fraction, lower = closer)")
|
||||
print(
|
||||
f"mean |err| vs fp32 (double) : {err_double.mean().item():.6e} (lower = more accurate)"
|
||||
)
|
||||
print(
|
||||
f"mean |err| vs fp32 (single) : {err_single.mean().item():.6e} (lower = more accurate)"
|
||||
)
|
||||
winner = (
|
||||
"single (direct fp32->fp8)"
|
||||
if err_single.mean() <= err_double.mean()
|
||||
else "double (via bf16)"
|
||||
)
|
||||
print(f"more accurate on average : {winner}")
|
||||
print(
|
||||
"NOTE: the born-fp8 kernel keeps the DOUBLE round on purpose to match "
|
||||
"the default path's rounding stages."
|
||||
)
|
||||
|
||||
|
||||
def run_config(
|
||||
name,
|
||||
num_tokens,
|
||||
num_heads,
|
||||
k_nope,
|
||||
iters,
|
||||
warmup,
|
||||
device,
|
||||
seed,
|
||||
magnitude,
|
||||
variants,
|
||||
sweep,
|
||||
):
|
||||
print(f"\n=== {name}: tokens={num_tokens} heads={num_heads} K={k_nope} ===")
|
||||
print(
|
||||
f" (K={k_nope} nope-in, N={N_LORA} nope-out, R={R_ROPE} rope; "
|
||||
"us/call LOWER = FASTER; GB/s HIGHER = BETTER; speedup >1 = new faster)"
|
||||
)
|
||||
q, q_nope, q_rope, w_kc = make_inputs(
|
||||
num_tokens, num_heads, k_nope, device, seed, magnitude
|
||||
)
|
||||
q_fp8 = torch.zeros(
|
||||
(num_tokens, num_heads, N_LORA + R_ROPE),
|
||||
dtype=torch.float8_e4m3fn,
|
||||
device=device,
|
||||
)
|
||||
|
||||
t_old = time_fn(
|
||||
lambda: old_path(q_fp8, q_nope, w_kc, q_rope, num_heads), iters, warmup
|
||||
)
|
||||
t_bmm = time_fn(lambda: old_path_bmm_only(q_nope, w_kc), iters, warmup)
|
||||
# standalone concat-cast (reads a fresh bf16 bmm out, like production)
|
||||
q_nope_out = torch.bmm(q_nope.transpose(0, 1), w_kc).transpose(0, 1)
|
||||
t_cast = time_fn(
|
||||
lambda: concat_and_cast_q_fp8_pad(q_fp8, q_nope_out, q_rope, num_heads),
|
||||
iters,
|
||||
warmup,
|
||||
)
|
||||
|
||||
b_old, b_new = analytic_bytes(num_tokens, num_heads, k_nope)
|
||||
print(
|
||||
f"OLD bmm (cublas bf16) : {t_bmm:10.1f} us/call"
|
||||
f" (component of OLD total)"
|
||||
)
|
||||
print(f"OLD concat_and_cast_q_fp8_pad: {t_cast:10.1f} us/call (component)")
|
||||
print(
|
||||
f"OLD total (bmm + concat-cast): {t_old:10.1f} us/call"
|
||||
f" ({b_old / 1e6:8.1f} MB analytic, {b_old / t_old / 1e3:7.0f} GB/s)"
|
||||
)
|
||||
|
||||
# Power-of-2 K collapses every variant to the same single-dot codegen.
|
||||
pow2 = k_nope & (k_nope - 1) == 0
|
||||
run_variants = ["auto"] if pow2 else variants
|
||||
ctx = make_check_ctx(num_tokens, num_heads, k_nope, device, seed + 1, magnitude)
|
||||
results = {}
|
||||
for v in run_variants:
|
||||
kw = {"variant": v}
|
||||
try:
|
||||
t_new = time_fn(
|
||||
lambda: new_path(q_fp8, q_nope, w_kc, q_rope, num_heads, **kw),
|
||||
iters,
|
||||
warmup,
|
||||
)
|
||||
except Exception as e:
|
||||
msg = (str(e).splitlines() or [type(e).__name__])[0]
|
||||
print(f"NEW {v:<24}: COMPILE/RUN FAIL — {msg[:100]}")
|
||||
continue
|
||||
c = check_variant(ctx, **kw)
|
||||
results[v] = t_new
|
||||
faster = "NEW FASTER" if t_new < t_old else "OLD FASTER"
|
||||
print(
|
||||
f"NEW {v:<24}: {t_new:10.1f} us/call"
|
||||
f" ({b_new / 1e6:8.1f} MB analytic, {b_new / t_new / 1e3:7.0f} GB/s,"
|
||||
f" speedup {t_old / t_new:5.2f}x vs OLD, {faster})"
|
||||
)
|
||||
rope = "PASS (bitwise identical)" if c["rope_bitexact"] else "FAIL (BUG)"
|
||||
print(
|
||||
f" rope bit-exact: {rope}; nope bitwise match vs OLD"
|
||||
f" {c['nope_bitwise_match_frac'] * 100:9.4f}% (100% = bit-exact);"
|
||||
f" max |dequant diff| {c['nope_max_dequant_absdiff']:.4f}"
|
||||
)
|
||||
print(
|
||||
f" nope |err| vs fp64 ref: mean old {c['nope_meanerr_old_vs_fp64']:.3e}"
|
||||
f" / new {c['nope_meanerr_new_vs_fp64']:.3e}; max old"
|
||||
f" {c['nope_maxerr_old_vs_fp64']:.3e} / new"
|
||||
f" {c['nope_maxerr_new_vs_fp64']:.3e} (lower = more accurate)"
|
||||
)
|
||||
if results:
|
||||
best = min(results, key=results.get)
|
||||
print(
|
||||
f"BEST variant : {best} @ {results[best]:.1f} us/call"
|
||||
f" (speedup {t_old / results[best]:.2f}x vs OLD total)"
|
||||
)
|
||||
|
||||
if sweep and results:
|
||||
print(
|
||||
f"\n--- tile sweep: {name} (us/call LOWER = FASTER;"
|
||||
" stages=0 -> Triton default) ---"
|
||||
)
|
||||
rows = []
|
||||
for v in results:
|
||||
for bm, bn, nw, ns in SWEEP_TILES:
|
||||
kw = dict(
|
||||
variant=v, block_m=bm, block_n=bn, num_warps=nw, num_stages=ns
|
||||
)
|
||||
try:
|
||||
t = time_fn(
|
||||
lambda: new_path(q_fp8, q_nope, w_kc, q_rope, num_heads, **kw),
|
||||
max(iters // 2, 20),
|
||||
warmup,
|
||||
)
|
||||
except Exception as e:
|
||||
msg = (str(e).splitlines() or [type(e).__name__])[0]
|
||||
print(
|
||||
f" {v:<10} bm={bm:<3} bn={bn:<3} warps={nw} stages={ns}:"
|
||||
f" FAIL — {msg[:70]}"
|
||||
)
|
||||
continue
|
||||
rows.append((t, v, bm, bn, nw, ns))
|
||||
print(
|
||||
f" {v:<10} bm={bm:<3} bn={bn:<3} warps={nw} stages={ns}:"
|
||||
f" {t:8.1f} us/call ({b_new / t / 1e3:5.0f} GB/s,"
|
||||
f" {t_old / t:5.2f}x vs OLD)"
|
||||
)
|
||||
rows.sort()
|
||||
print(" -- top 5 (fastest first) --")
|
||||
for t, v, bm, bn, nw, ns in rows[:5]:
|
||||
print(
|
||||
f" {v:<10} bm={bm:<3} bn={bn:<3} warps={nw} stages={ns}:"
|
||||
f" {t:8.1f} us/call ({t_old / t:5.2f}x vs OLD)"
|
||||
)
|
||||
|
||||
|
||||
def main():
|
||||
parser = argparse.ArgumentParser(description=__doc__)
|
||||
parser.add_argument("--tokens", type=int, default=4096, help="s_q per call")
|
||||
parser.add_argument("--iters", type=int, default=200)
|
||||
parser.add_argument("--warmup", type=int, default=20)
|
||||
parser.add_argument("--seed", type=int, default=1234)
|
||||
parser.add_argument(
|
||||
"--magnitude",
|
||||
type=float,
|
||||
default=1.0,
|
||||
help="input scale multiplier (q amax stress)",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--heads",
|
||||
type=int,
|
||||
default=None,
|
||||
help="run a single head count instead of the GLM(64,K192)+DS(128,K128) pair",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--k-nope",
|
||||
type=int,
|
||||
default=128,
|
||||
help="qk_nope_head_dim for --heads runs (ignored for the default pair)",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--variants",
|
||||
type=str,
|
||||
default="all",
|
||||
help=(
|
||||
"comma list of non-power-of-2-K variants to bench "
|
||||
f"(default: all = {','.join(ALL_VARIANTS)}); power-of-2-K configs "
|
||||
"always run the single collapsed 'auto' variant"
|
||||
),
|
||||
)
|
||||
parser.add_argument(
|
||||
"--sweep",
|
||||
action="store_true",
|
||||
help="also sweep (block_m, block_n, num_warps, num_stages) per variant",
|
||||
)
|
||||
parser.add_argument("--rounding-study", action="store_true")
|
||||
args = parser.parse_args()
|
||||
|
||||
if args.variants == "all":
|
||||
variants = ALL_VARIANTS
|
||||
else:
|
||||
variants = [v.strip() for v in args.variants.split(",") if v.strip()]
|
||||
unknown = set(variants) - set(ALL_VARIANTS) - {"auto"}
|
||||
assert not unknown, f"unknown variants: {sorted(unknown)}"
|
||||
|
||||
assert torch.cuda.is_available(), "CUDA GPU required"
|
||||
device = torch.device("cuda")
|
||||
name = torch.cuda.get_device_name(device)
|
||||
print(f"device: {name}; torch {torch.__version__}")
|
||||
|
||||
if args.heads is not None:
|
||||
configs = [(f"custom h={args.heads} K={args.k_nope}", args.heads, args.k_nope)]
|
||||
else:
|
||||
configs = [
|
||||
("GLM-5.2 (h=64, K=192)", 64, 192),
|
||||
("DS-V3.2 (h=128, K=128)", 128, 128),
|
||||
]
|
||||
for cfg_name, heads, k_nope in configs:
|
||||
run_config(
|
||||
cfg_name,
|
||||
args.tokens,
|
||||
heads,
|
||||
k_nope,
|
||||
args.iters,
|
||||
args.warmup,
|
||||
device,
|
||||
args.seed,
|
||||
args.magnitude,
|
||||
variants,
|
||||
args.sweep,
|
||||
)
|
||||
|
||||
if args.rounding_study:
|
||||
rounding_study(device, args.seed)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -105,7 +105,7 @@ import { Playground } from "/src/snippets/_playground.jsx";
|
||||
|
||||
## 2. Configuration Tips
|
||||
|
||||
- **DeepSeek Sparse Attention (DSA).** GLM-5.2 uses the `glm_moe_dsa` architecture; SGLang auto-selects the DSA attention backends (`flashmla_sparse` prefill, `fa3` decode, `sgl-kernel` indexer topk). No attention-backend flag is needed on the supported hardware. SGLang also auto-selects the KV-cache dtype for DSA models — `fp8_e4m3` on Blackwell (B200/GB300/B300, which then routes DSA through the TensorRT-LLM backend) and `bf16` on Hopper (H200) — so no `--kv-cache-dtype` flag is required.
|
||||
- **DeepSeek Sparse Attention (DSA).** GLM-5.2 uses the `glm_moe_dsa` architecture; SGLang auto-selects the DSA attention backends (`flashmla_sparse` prefill, `fa3` decode, `sgl-kernel` indexer topk). No attention-backend flag is needed on the supported hardware. SGLang also auto-selects the KV-cache dtype for DSA models — `fp8_e4m3` on Blackwell (B200/GB300/B300, which then routes DSA through the TensorRT-LLM backend) and `bf16` on Hopper (H200) — so no `--kv-cache-dtype` flag is required. On Hopper, pairing `--kv-cache-dtype fp8_e4m3` with `--dsa-prefill-backend flashmla_sparse_q8 --dsa-decode-backend flashmla_kv` selects the native FP8 sparse prefill kernel (computes directly on the fp8 KV cache with no fp8→bf16 dequantization round-trip; GLM-5.2's 64 query heads match the kernel's native tile) — see the [DeepSeek-V3.2 page](../DeepSeek/DeepSeek-V3_2) for kernel details; the optional `SGLANG_ENABLE_DSA_Q8KV8_*` performance env vars are documented in `python/sglang/srt/environ.py`.
|
||||
- **MTP / speculative decoding.** The checkpoint ships one nextn layer. Enable EAGLE MTP for lower latency (`--speculative-algorithm EAGLE --speculative-num-steps 5 --speculative-eagle-topk 1 --speculative-num-draft-tokens 6` for low-latency; `1-1-2` for balanced). The config's `index_share_for_mtp_iteration` reuses the DSA indexer's topk across draft steps (effective only at `--speculative-eagle-topk 1`). **Tune the draft length to the accept length.** GLM-5.2's MTP head is strong — accept length runs high (4+ in many workloads, near-saturating at 5–6 in low-latency runs). Watch the server's reported **accept length** and adjust `--speculative-num-steps` / `--speculative-num-draft-tokens` accordingly: while accept length stays close to the draft-token count there is headroom to push them higher (more accepted tokens per step); if it falls well below, lower them — every rejected draft token is wasted verification compute.
|
||||
- **Memory.** The FP8 weights are large (MoE total, not active params). Start around `--mem-fraction-static 0.8` on H200 (TP8) and tune up; raise it for the 4-GPU GB300 single-node layout (TP4).
|
||||
- **DP-Attention + DeepEP** for the balanced/high-throughput strategies spreads attention across data-parallel ranks and routes MoE through DeepEP.
|
||||
|
||||
@@ -0,0 +1,84 @@
|
||||
/* Copyright 2026 SGLang Team. All Rights Reserved.
|
||||
|
||||
Licensed under the Apache License, Version 2.0 (the "License");
|
||||
you may not use this file except in compliance with the License.
|
||||
You may obtain a copy of the License at
|
||||
|
||||
http://www.apache.org/licenses/LICENSE-2.0
|
||||
|
||||
Unless required by applicable law or agreed to in writing, software
|
||||
distributed under the License is distributed on an "AS IS" BASIS,
|
||||
WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
See the License for the specific language governing permissions and
|
||||
limitations under the License.
|
||||
==============================================================================*/
|
||||
|
||||
// JIT dispatch entry for the SM90 Q8KV8 born-fp8 q-prep kernel.
|
||||
#pragma once
|
||||
|
||||
#include <dlpack/dlpack.h>
|
||||
#include <tvm/ffi/container/tensor.h>
|
||||
|
||||
#include "kernel.cuh"
|
||||
#include <cstdint>
|
||||
#include <cuda_runtime.h>
|
||||
|
||||
namespace {
|
||||
|
||||
// All strides are in elements; validation of dtypes/shapes/alignment happens
|
||||
// in the Python wrapper (sglang/kernels/ops/attention/qprep_bf16_fp8_sm90.py).
|
||||
void qprep_bf16_fp8_dispatch(
|
||||
tvm::ffi::TensorView q_nope,
|
||||
tvm::ffi::TensorView w_kc,
|
||||
tvm::ffi::TensorView q_rope,
|
||||
tvm::ffi::TensorView out,
|
||||
int64_t num_tokens,
|
||||
int64_t num_heads,
|
||||
int64_t k_dim,
|
||||
int64_t a_s0,
|
||||
int64_t a_s1,
|
||||
int64_t b_s0,
|
||||
int64_t b_s2,
|
||||
int64_t r_s0,
|
||||
int64_t r_s1,
|
||||
int64_t o_s0,
|
||||
int64_t o_s1,
|
||||
int64_t rope_vec16,
|
||||
int64_t out_vec16,
|
||||
int64_t cuda_stream) {
|
||||
QprepBf16Fp8Sm90Params params;
|
||||
params.num_tokens = (int)num_tokens;
|
||||
params.num_heads = (int)num_heads;
|
||||
params.q_nope = q_nope.data_ptr();
|
||||
params.a_s0 = a_s0;
|
||||
params.a_s1 = a_s1;
|
||||
params.w_kc = w_kc.data_ptr();
|
||||
params.b_s0 = b_s0;
|
||||
params.b_s2 = b_s2;
|
||||
params.q_rope = q_rope.data_ptr();
|
||||
params.r_s0 = r_s0;
|
||||
params.r_s1 = r_s1;
|
||||
params.rope_vec16 = (bool)rope_vec16;
|
||||
params.out = out.data_ptr();
|
||||
params.o_s0 = o_s0;
|
||||
params.o_s1 = o_s1;
|
||||
params.out_vec16 = (bool)out_vec16;
|
||||
|
||||
DLDevice dev = q_nope.device();
|
||||
cudaSetDevice(dev.device_id);
|
||||
params.stream = reinterpret_cast<cudaStream_t>(cuda_stream);
|
||||
|
||||
switch (k_dim) {
|
||||
case 128:
|
||||
qprep_sm90::run_qprep_bf16_fp8_sm90<128>(params);
|
||||
return;
|
||||
case 192:
|
||||
qprep_sm90::run_qprep_bf16_fp8_sm90<192>(params);
|
||||
return;
|
||||
default:
|
||||
fprintf(stderr, "qprep_bf16_fp8_sm90: unsupported k_dim=%ld (must be 128 or 192)\n", (long)k_dim);
|
||||
exit(1);
|
||||
}
|
||||
}
|
||||
|
||||
} // namespace
|
||||
@@ -0,0 +1,516 @@
|
||||
/* Copyright 2026 SGLang Team. All Rights Reserved.
|
||||
|
||||
Licensed under the Apache License, Version 2.0 (the "License");
|
||||
you may not use this file except in compliance with the License.
|
||||
You may obtain a copy of the License at
|
||||
|
||||
http://www.apache.org/licenses/LICENSE-2.0
|
||||
|
||||
Unless required by applicable law or agreed to in writing, software
|
||||
distributed under the License is distributed on an "AS IS" BASIS,
|
||||
WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
See the License for the specific language governing permissions and
|
||||
limitations under the License.
|
||||
==============================================================================*/
|
||||
|
||||
// SM90 (Hopper) Q8KV8 born-fp8 q-prep kernel.
|
||||
//
|
||||
// Computes, per head h:
|
||||
// out[:, h, :512] = fp8_e4m3(bf16(fp32_accum(q_nope[:, h, :] @ w_kc[h])))
|
||||
// out[:, h, 512:576] = fp8_e4m3(q_rope[:, h, :])
|
||||
//
|
||||
// This is the CUDA replacement for the Triton absorbed_bmm_concat_cast_q_fp8
|
||||
// kernel (triton_ops/cache_ops.py). The epilogue keeps the exact rounding
|
||||
// chain of the Triton variants: fp32 WGMMA accumulate -> bf16 round-to-nearest
|
||||
// (cublas-equivalent output rounding) -> fp8_e4m3 rn/satfinite on store. The
|
||||
// K dimension is consumed as one in-order chain of k=16 WGMMA steps into a
|
||||
// single fp32 accumulator, i.e. the same fp32 add order as the Triton
|
||||
// "two_dot"/"grouped" variants (128+64 chained tl.dot), so the nope half can
|
||||
// come out bitwise identical to them.
|
||||
//
|
||||
// Phase-2 design (2 CTAs/SM + double-buffered B):
|
||||
// grid = (ceil(T / 128), H); one CTA = two WGMMA warpgroups (256 threads)
|
||||
// owning a 128-row m-tile of one head (warpgroup w computes rows
|
||||
// [64w, 64w+64)). The A tile [128, K] bf16 is cp.async'd to smem once (L2
|
||||
// evict_first: streamed) and the rope path runs under that load's wait.
|
||||
// The N=512 output is produced in N_SLABS n-slabs of BN columns; the B
|
||||
// slab [BN, K] bf16 is double-buffered (L2 evict_last: re-read by every
|
||||
// CTA of the head) and prefetched one full round ahead. Per round, the
|
||||
// fp8 stage-write -> barrier -> refill-issue -> coalesced-flush order
|
||||
// makes one barrier serve both the stage handoff and the CTA-wide WGMMA
|
||||
// drain of the buffer being refilled, and the flush plus the next round's
|
||||
// gemm overlap the refill. BN is sized so that A + 2 B buffers + the fp8
|
||||
// stage fit in half an SM's smem, keeping 2 CTAs co-resident per SM
|
||||
// (register cap 128 via launch bounds; measured faster than every
|
||||
// 1-CTA/SM variant tried, including wider CTAs and dual-accumulator
|
||||
// cross-round software pipelines): K=192 -> BN=64 (104 KB), K=128 ->
|
||||
// BN=128 (112 KB). The round loop is left un-unrolled when N_SLABS > 4:
|
||||
// full unrolling blows the 128-register budget and spills to local.
|
||||
|
||||
#pragma once
|
||||
|
||||
#include <cute/tensor.hpp>
|
||||
#include <cutlass/bfloat16.h>
|
||||
|
||||
#include "params.h"
|
||||
#include <cstdint>
|
||||
#include <cstdio>
|
||||
#include <cstdlib>
|
||||
#include <cuda_bf16.h>
|
||||
#include <type_traits>
|
||||
|
||||
namespace qprep_sm90 {
|
||||
|
||||
using namespace cute;
|
||||
using bf16 = cutlass::bfloat16_t;
|
||||
|
||||
#define QPREP_ASSERT(cond) \
|
||||
do { \
|
||||
if (!(cond)) { \
|
||||
fprintf(stderr, "QPREP_ASSERT failed (%s:%d): %s\n", __FILE__, __LINE__, #cond); \
|
||||
exit(1); \
|
||||
} \
|
||||
} while (0)
|
||||
|
||||
#define QPREP_CUDA_CHECK(call) \
|
||||
do { \
|
||||
cudaError_t err = (call); \
|
||||
if (err != cudaSuccess) { \
|
||||
fprintf(stderr, "CUDA error (%s:%d): %s\n", __FILE__, __LINE__, cudaGetErrorString(err)); \
|
||||
exit(1); \
|
||||
} \
|
||||
} while (0)
|
||||
|
||||
__host__ __device__ __forceinline__ constexpr int ceil_div_i(int a, int b) {
|
||||
return (a + b - 1) / b;
|
||||
}
|
||||
|
||||
// ---------------------------------------------------------------------------
|
||||
// Device helpers
|
||||
// ---------------------------------------------------------------------------
|
||||
|
||||
// L2 eviction policies (same helpers as the sparse-prefill kernel): A/rope
|
||||
// are streamed once (evict_first); the per-head w_kc slice is re-read from L2
|
||||
// by every CTA of the head (evict_last).
|
||||
__device__ __forceinline__ int64_t createpolicy_evict_last() {
|
||||
int64_t res;
|
||||
asm volatile("createpolicy.fractional.L2::evict_last.b64 %0, 1.0; \n\t" : "=l"(res) :);
|
||||
return res;
|
||||
}
|
||||
|
||||
__device__ __forceinline__ int64_t createpolicy_evict_first() {
|
||||
int64_t res;
|
||||
asm volatile("createpolicy.fractional.L2::evict_first.b64 %0, 1.0; \n\t" : "=l"(res) :);
|
||||
return res;
|
||||
}
|
||||
|
||||
// 16-byte cp.async.cg with an L2 cache policy, zero-filling when pred is
|
||||
// false (same instruction family as the sparse-prefill producer).
|
||||
__device__ __forceinline__ void
|
||||
cp_async_16_zfill(void* smem_dst, const void* gmem_src, bool pred, int64_t cache_policy) {
|
||||
uint32_t dst_addr = cute::cast_smem_ptr_to_uint(smem_dst);
|
||||
asm volatile(
|
||||
"cp.async.cg.shared.global.L2::cache_hint.L2::256B [%0], [%1], 16, %2, %3;\n" ::"r"(dst_addr),
|
||||
"l"(gmem_src),
|
||||
"r"(pred ? 16 : 0),
|
||||
"l"(cache_policy));
|
||||
}
|
||||
|
||||
__device__ __forceinline__ void cp_async_16(void* smem_dst, const void* gmem_src, int64_t cache_policy) {
|
||||
uint32_t dst_addr = cute::cast_smem_ptr_to_uint(smem_dst);
|
||||
asm volatile(
|
||||
"cp.async.cg.shared.global.L2::cache_hint.L2::256B [%0], [%1], 16, %2;\n" ::"r"(dst_addr),
|
||||
"l"(gmem_src),
|
||||
"l"(cache_policy));
|
||||
}
|
||||
|
||||
// Pack two fp32 into two fp8_e4m3 bytes with round-to-nearest + satfinite.
|
||||
// PTX: cvt.rn.satfinite.e4m3x2.f32 d, a, b -> d[7:0] = cvt(b), d[15:8] = cvt(a).
|
||||
__device__ __forceinline__ uint16_t f32x2_to_e4m3x2_rn_satfinite(float f_lo, float f_hi) {
|
||||
uint16_t v;
|
||||
asm volatile("cvt.rn.satfinite.e4m3x2.f32 %0, %1, %2;\n" : "=h"(v) : "f"(f_hi), "f"(f_lo));
|
||||
return v;
|
||||
}
|
||||
|
||||
// The exact Triton epilogue rounding chain for the nope half:
|
||||
// fp32 accum -> bf16 (rn) -> fp32 (exact) -> fp8_e4m3 (rn, satfinite).
|
||||
__device__ __forceinline__ uint16_t f32x2_to_bf16x2_to_e4m3x2(float f0, float f1) {
|
||||
const __nv_bfloat162 b = __float22bfloat162_rn(make_float2(f0, f1));
|
||||
return f32x2_to_e4m3x2_rn_satfinite(__low2float(b), __high2float(b));
|
||||
}
|
||||
|
||||
// ---------------------------------------------------------------------------
|
||||
// Kernel
|
||||
// ---------------------------------------------------------------------------
|
||||
|
||||
template <typename Kernel>
|
||||
__global__ void qprep_bf16_fp8_kernel(__grid_constant__ const QprepBf16Fp8Sm90Params params);
|
||||
|
||||
template <int K_DIM>
|
||||
struct QprepBf16Fp8Kernel {
|
||||
static constexpr int BM = 128; // m-tile rows (two WGMMA warpgroups)
|
||||
// n-slab width: sized so that A + 2 B buffers + the fp8 stage fit in half
|
||||
// an SM's smem -> 2 CTAs/SM (measured worth more than any intra-CTA
|
||||
// pipelining): K=128 fits BN=128 (112 KB); K=192 needs BN=64 (104 KB).
|
||||
static constexpr int BN = (K_DIM > 128) ? 64 : 128;
|
||||
static constexpr int N_OUT = 512; // kv_lora_rank
|
||||
static constexpr int ROPE = 64; // qk_rope_head_dim
|
||||
static constexpr int NUM_THREADS = 256;
|
||||
static constexpr int N_SLABS = N_OUT / BN;
|
||||
static constexpr int LOAD_ROWS_PER_PASS = NUM_THREADS / 8; // 16B-chunk loaders
|
||||
// 2 CTAs/SM co-residency (register cap 128 via launch bounds). Measured
|
||||
// faster than every 1-CTA/SM variant tried (wider CTAs, dual-accumulator
|
||||
// cross-round software pipelines).
|
||||
static constexpr int MIN_CTAS = 2;
|
||||
|
||||
static_assert(K_DIM % 64 == 0, "K must tile the SW128 bf16 GMMA atom (64 cols)");
|
||||
static_assert(N_OUT % BN == 0);
|
||||
|
||||
// K-major SW128 smem layouts for the SS WGMMA operands (bf16 atom = 8x64).
|
||||
using SmemLayoutA =
|
||||
decltype(tile_to_shape(GMMA::Layout_K_SW128_Atom<bf16>{}, Shape<Int<BM>, Int<K_DIM>>{}, Step<_1, _2>{}));
|
||||
using SmemLayoutB =
|
||||
decltype(tile_to_shape(GMMA::Layout_K_SW128_Atom<bf16>{}, Shape<Int<BN>, Int<K_DIM>>{}, Step<_1, _2>{}));
|
||||
|
||||
// Two warpgroups stacked along M: threads [128w, 128w+128) own rows
|
||||
// [64w, 64w+64) of the m-tile. The atom's N width must match BN.
|
||||
using MmaAtom_t = std::conditional_t<
|
||||
BN == 128,
|
||||
SM90_64x128x16_F32BF16BF16_SS<GMMA::Major::K, GMMA::Major::K>,
|
||||
SM90_64x64x16_F32BF16BF16_SS<GMMA::Major::K, GMMA::Major::K>>;
|
||||
using TiledMMA_t = decltype(make_tiled_mma(MmaAtom_t{}, Layout<Shape<_2, _1, _1>>{}));
|
||||
|
||||
struct SharedStorage {
|
||||
array_aligned<bf16, cosize_v<SmemLayoutA>, 128> a; // resident A m-block
|
||||
array_aligned<bf16, cosize_v<SmemLayoutB>, 128> b[2]; // double-buffered B slab
|
||||
// fp8 output staging for one n-slab: scattered per-thread u16 epilogue
|
||||
// writes land here, then leave as coalesced 16B global stores (the direct
|
||||
// u16 global stores 4x-amplify the store sectors and throttle the LSU).
|
||||
// Single buffer: the end-of-round B wait barrier separates one round's
|
||||
// copy-out reads from the next round's stage writes.
|
||||
array_aligned<uint8_t, BM * BN, 16> c_stage;
|
||||
};
|
||||
|
||||
// -------------------------------------------------------------------------
|
||||
// Loads: NUM_THREADS as (NUM_THREADS/8) row-threads x 8 chunk-threads, 16B
|
||||
// per cp.async. Smem addresses go through the CUTE tensor so the SW128
|
||||
// swizzle is applied (16B chunks stay contiguous under the swizzle).
|
||||
// -------------------------------------------------------------------------
|
||||
template <typename SmemT>
|
||||
static __device__ __forceinline__ void
|
||||
load_a_tile(SmemT& sA, const bf16* gA, int64_t a_s0, int m_residue, int tid, int64_t cache_policy) {
|
||||
const int cthr = tid % 8, rthr = tid / 8;
|
||||
CUTE_UNROLL
|
||||
for (int mi = 0; mi < BM / LOAD_ROWS_PER_PASS; ++mi) {
|
||||
const int row = rthr + LOAD_ROWS_PER_PASS * mi;
|
||||
const bool pred = row < m_residue; // zfill OOB rows: 0 * w == 0, never stored
|
||||
const bf16* g = gA + (int64_t)row * a_s0;
|
||||
CUTE_UNROLL
|
||||
for (int ki = 0; ki < K_DIM / 64; ++ki) {
|
||||
const int col = cthr * 8 + 64 * ki;
|
||||
cp_async_16_zfill(&sA(row, col), g + col, pred, cache_policy);
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
template <typename SmemT>
|
||||
static __device__ __forceinline__ void
|
||||
load_b_slab(SmemT& sB, const bf16* gB_head, int64_t b_s2, int nb, int tid, int64_t cache_policy) {
|
||||
const int cthr = tid % 8, rthr = tid / 8;
|
||||
const bf16* g0 = gB_head + (int64_t)nb * BN * b_s2;
|
||||
CUTE_UNROLL
|
||||
for (int ni = 0; ni < BN / LOAD_ROWS_PER_PASS; ++ni) {
|
||||
const int nrow = rthr + LOAD_ROWS_PER_PASS * ni;
|
||||
const bf16* g = g0 + (int64_t)nrow * b_s2;
|
||||
CUTE_UNROLL
|
||||
for (int ki = 0; ki < K_DIM / 64; ++ki) {
|
||||
const int col = cthr * 8 + 64 * ki;
|
||||
cp_async_16(&sB(nrow, col), g + col, cache_policy);
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// -------------------------------------------------------------------------
|
||||
// SS WGMMA over the whole K extent as one in-order k=16 chain (clears the
|
||||
// accumulator on the first step). Adapted from the sparse-prefill gemm_ss.
|
||||
// -------------------------------------------------------------------------
|
||||
template <typename TA, typename TB, typename TC>
|
||||
static __device__ __forceinline__ void gemm_ss(TiledMMA_t& tiled_mma, TA const& sA, TB const& sB, TC& acc, int tid) {
|
||||
ThrMMA thr_mma = tiled_mma.get_slice(tid);
|
||||
Tensor sA_frag = thr_mma.partition_fragment_A(sA);
|
||||
Tensor sB_frag = thr_mma.partition_fragment_B(sB);
|
||||
static_assert(size<2>(sA_frag) == size<2>(sB_frag));
|
||||
|
||||
warpgroup_fence_operand(acc);
|
||||
warpgroup_arrive();
|
||||
tiled_mma.accumulate_ = GMMA::ScaleOut::Zero;
|
||||
CUTE_UNROLL
|
||||
for (int k = 0; k < size<2>(sA_frag); ++k) {
|
||||
cute::gemm(tiled_mma, sA_frag(_, _, k), sB_frag(_, _, k), acc);
|
||||
tiled_mma.accumulate_ = GMMA::ScaleOut::One;
|
||||
}
|
||||
warpgroup_fence_operand(acc);
|
||||
}
|
||||
|
||||
// -------------------------------------------------------------------------
|
||||
// Epilogue for one n-slab: fp32 acc -> bf16 -> fp8, 2 adjacent columns per
|
||||
// 16-bit store. WGMMA m64nN C layout: within its warpgroup, thread t holds
|
||||
// rows (t/32)*16 + (t%32)/4 + {0,8} (plus 64 * warpgroup_idx here) and
|
||||
// columns (t%4)*2 + 8j + {0,1}; fragment linear index
|
||||
// i = 4j + 2*row_parity + col_parity.
|
||||
// -------------------------------------------------------------------------
|
||||
template <typename TC>
|
||||
static __device__ __forceinline__ void
|
||||
store_slab_direct(TC const& acc, uint8_t* gO, int64_t o_s0, int n0, int row_base, int col_base, int m_residue) {
|
||||
CUTE_UNROLL
|
||||
for (int rp = 0; rp < 2; ++rp) {
|
||||
const int row = row_base + 8 * rp;
|
||||
if (row >= m_residue) continue;
|
||||
uint8_t* orow = gO + (int64_t)row * o_s0 + n0 + col_base;
|
||||
CUTE_UNROLL
|
||||
for (int j = 0; j < BN / 8; ++j) {
|
||||
const float f0 = acc(j * 4 + rp * 2 + 0);
|
||||
const float f1 = acc(j * 4 + rp * 2 + 1);
|
||||
*reinterpret_cast<uint16_t*>(orow + 8 * j) = f32x2_to_bf16x2_to_e4m3x2(f0, f1);
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// Staged variant: XOR-swizzle the 16B chunk index by the row so the u16
|
||||
// stage writes (8 distinct rows per warp) spread across banks, while the
|
||||
// 16B copy-out reads stay conflict-free row segments. The XOR must be
|
||||
// masked to the chunks actually present in a BN-wide row.
|
||||
static constexpr int STAGE_CHUNK_MASK = BN / 16 - 1;
|
||||
static __device__ __forceinline__ int stage_off(int r, int c) {
|
||||
const int phys = ((c >> 4) ^ r) & STAGE_CHUNK_MASK;
|
||||
return r * BN + (phys << 4) + (c & 15);
|
||||
}
|
||||
|
||||
// Stage-write half: fp32 acc -> bf16 -> fp8 u16 writes into the swizzled
|
||||
// smem stage. Reads only the accumulator, so it can run while the next
|
||||
// round's WGMMA chain and the B refill are in flight. The caller provides
|
||||
// the __syncthreads() handoff before stage_flush.
|
||||
template <typename TC>
|
||||
static __device__ __forceinline__ void stage_write(TC const& acc, uint8_t* stage, int row_base, int col_base) {
|
||||
CUTE_UNROLL
|
||||
for (int rp = 0; rp < 2; ++rp) {
|
||||
const int row = row_base + 8 * rp; // OOB rows staged but never copied out
|
||||
CUTE_UNROLL
|
||||
for (int j = 0; j < BN / 8; ++j) {
|
||||
const float f0 = acc(j * 4 + rp * 2 + 0);
|
||||
const float f1 = acc(j * 4 + rp * 2 + 1);
|
||||
*reinterpret_cast<uint16_t*>(stage + stage_off(row, col_base + 8 * j)) = f32x2_to_bf16x2_to_e4m3x2(f0, f1);
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// Copy-out half: coalesced 16B stores of the staged fp8 slab.
|
||||
static __device__ __forceinline__ void
|
||||
stage_flush(const uint8_t* stage, uint8_t* gO, int64_t o_s0, int n0, int m_residue, int tid) {
|
||||
constexpr int CHUNKS_PER_ROW = BN / 16;
|
||||
constexpr int NUM_CHUNKS = BM * BN / 16;
|
||||
CUTE_UNROLL
|
||||
for (int i = 0; i < NUM_CHUNKS / NUM_THREADS; ++i) {
|
||||
const int chunk = tid + i * NUM_THREADS;
|
||||
const int r = chunk / CHUNKS_PER_ROW;
|
||||
const int c = (chunk % CHUNKS_PER_ROW) * 16;
|
||||
if (r >= m_residue) continue;
|
||||
const uint4 v = *reinterpret_cast<const uint4*>(stage + stage_off(r, c));
|
||||
*reinterpret_cast<uint4*>(gO + (int64_t)r * o_s0 + n0 + c) = v;
|
||||
}
|
||||
}
|
||||
|
||||
// -------------------------------------------------------------------------
|
||||
// Rope path: out[:, h, 512:576] = fp8(q_rope[:, h, :]). bf16 -> fp32
|
||||
// (exact) -> fp8 rn/satfinite == the Triton store conversion, so this half
|
||||
// is bit-exact vs concat_and_cast_q_fp8_pad. 8 bf16 per thread-chunk.
|
||||
// -------------------------------------------------------------------------
|
||||
static __device__ __forceinline__ void
|
||||
rope_path(const bf16* gR, uint8_t* gO, const QprepBf16Fp8Sm90Params& p, int m_residue, int tid) {
|
||||
const int cthr = tid % 8, rthr = tid / 8;
|
||||
constexpr int PASSES = BM / LOAD_ROWS_PER_PASS;
|
||||
if (p.rope_vec16 && p.out_vec16) {
|
||||
// Fast path: batch-issue every row's uint4 load first so the load
|
||||
// latencies pipeline (one exposed latency instead of PASSES chained
|
||||
// load-use stalls), then convert + store.
|
||||
uint4 raw[PASSES];
|
||||
CUTE_UNROLL
|
||||
for (int mi = 0; mi < PASSES; ++mi) {
|
||||
const int row = rthr + LOAD_ROWS_PER_PASS * mi;
|
||||
if (row >= m_residue) continue;
|
||||
raw[mi] = *reinterpret_cast<const uint4*>(gR + (int64_t)row * p.r_s0 + cthr * 8);
|
||||
}
|
||||
CUTE_UNROLL
|
||||
for (int mi = 0; mi < PASSES; ++mi) {
|
||||
const int row = rthr + LOAD_ROWS_PER_PASS * mi;
|
||||
if (row >= m_residue) continue;
|
||||
const uint32_t* w = reinterpret_cast<const uint32_t*>(&raw[mi]);
|
||||
uint16_t packed[4];
|
||||
CUTE_UNROLL
|
||||
for (int i = 0; i < 4; ++i) {
|
||||
const __nv_bfloat162 v = *reinterpret_cast<const __nv_bfloat162*>(&w[i]);
|
||||
packed[i] = f32x2_to_e4m3x2_rn_satfinite(__low2float(v), __high2float(v));
|
||||
}
|
||||
// out_vec16 guarantees 16B-aligned rows; N_OUT + 8*cthr keeps 8B
|
||||
// alignment, so the 8-byte chunk goes out as one coalesced store.
|
||||
*reinterpret_cast<uint64_t*>(gO + (int64_t)row * p.o_s0 + N_OUT + cthr * 8) =
|
||||
*reinterpret_cast<const uint64_t*>(packed);
|
||||
}
|
||||
return;
|
||||
}
|
||||
// Unaligned fallback: element strides only guarantee 2B alignment.
|
||||
CUTE_UNROLL
|
||||
for (int mi = 0; mi < PASSES; ++mi) {
|
||||
const int row = rthr + LOAD_ROWS_PER_PASS * mi;
|
||||
if (row >= m_residue) continue;
|
||||
const bf16* g = gR + (int64_t)row * p.r_s0 + cthr * 8;
|
||||
uint8_t* o = gO + (int64_t)row * p.o_s0 + N_OUT + cthr * 8;
|
||||
const __nv_bfloat16* gh = reinterpret_cast<const __nv_bfloat16*>(g);
|
||||
CUTE_UNROLL
|
||||
for (int i = 0; i < 4; ++i) {
|
||||
const __nv_bfloat162 v = __nv_bfloat162(gh[2 * i], gh[2 * i + 1]);
|
||||
*reinterpret_cast<uint16_t*>(o + 2 * i) = f32x2_to_e4m3x2_rn_satfinite(__low2float(v), __high2float(v));
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// -------------------------------------------------------------------------
|
||||
// Main device function
|
||||
// -------------------------------------------------------------------------
|
||||
static __device__ __forceinline__ void devfunc(const QprepBf16Fp8Sm90Params& p) {
|
||||
#if defined(__CUDA_ARCH__) && (__CUDA_ARCH__ == 900)
|
||||
const int m0 = blockIdx.x * BM;
|
||||
const int h = blockIdx.y;
|
||||
const int tid = threadIdx.x;
|
||||
const int m_residue = p.num_tokens - m0; // > 0 by grid construction
|
||||
|
||||
extern __shared__ char smem_raw[];
|
||||
SharedStorage& ss = *reinterpret_cast<SharedStorage*>(smem_raw);
|
||||
Tensor sA = make_tensor(make_smem_ptr(ss.a.data()), SmemLayoutA{});
|
||||
Tensor sB0 = make_tensor(make_smem_ptr(ss.b[0].data()), SmemLayoutB{});
|
||||
Tensor sB1 = make_tensor(make_smem_ptr(ss.b[1].data()), SmemLayoutB{});
|
||||
|
||||
const bf16* gA = reinterpret_cast<const bf16*>(p.q_nope) + (int64_t)m0 * p.a_s0 + (int64_t)h * p.a_s1;
|
||||
const bf16* gB = reinterpret_cast<const bf16*>(p.w_kc) + (int64_t)h * p.b_s0;
|
||||
const bf16* gR = reinterpret_cast<const bf16*>(p.q_rope) + (int64_t)m0 * p.r_s0 + (int64_t)h * p.r_s1;
|
||||
uint8_t* gO = reinterpret_cast<uint8_t*>(p.out) + (int64_t)m0 * p.o_s0 + (int64_t)h * p.o_s1;
|
||||
|
||||
const int64_t policy_stream = createpolicy_evict_first();
|
||||
const int64_t policy_keep = createpolicy_evict_last();
|
||||
|
||||
// Issue the A block + B slab 0 (group 0), then B slab 1 (group 1), then
|
||||
// run the rope path over the in-flight async loads.
|
||||
load_a_tile(sA, gA, p.a_s0, m_residue, tid, policy_stream);
|
||||
load_b_slab(sB0, gB, p.b_s2, 0, tid, policy_keep);
|
||||
cp_async_fence();
|
||||
load_b_slab(sB1, gB, p.b_s2, 1, tid, policy_keep);
|
||||
cp_async_fence();
|
||||
|
||||
// Rope in the prologue: its global-load latency hides under the wait for
|
||||
// the A/B cp.async stream (measured better than placing it after the
|
||||
// first gemm commit on the 1-CTA/SM K=192 path).
|
||||
rope_path(gR, gO, p, m_residue, tid);
|
||||
|
||||
cp_async_wait<1>(); // A tile + B0 done; B1 still in flight
|
||||
__syncthreads();
|
||||
|
||||
TiledMMA_t tiled_mma;
|
||||
const int row_base = (tid / 128) * 64 + ((tid % 128) / 32) * 16 + ((tid % 32) / 4);
|
||||
const int col_base = (tid % 4) * 2;
|
||||
|
||||
{
|
||||
// Single accumulator: 2-CTA/SM co-residency covers the epilogue
|
||||
// latency (measured faster than every cross-round dual-accumulator
|
||||
// pipeline variant, which needs >128 regs and forfeits co-residency);
|
||||
// the B double-buffer still prefetches slab nb+1 a full round ahead.
|
||||
Tensor acc = partition_fragment_C(tiled_mma, Shape<Int<BM>, Int<BN>>{});
|
||||
gemm_ss(tiled_mma, sA, sB0, acc, tid);
|
||||
warpgroup_commit_batch();
|
||||
|
||||
auto round_body = [&](int nb) __attribute__((always_inline)) {
|
||||
warpgroup_wait<0>(); // gemm(nb) drained
|
||||
if (p.out_vec16) {
|
||||
// Single barrier: stage handoff + CTA-wide WGMMA drain of B[nb%2].
|
||||
stage_write(acc, ss.c_stage.data(), row_base, col_base);
|
||||
__syncthreads();
|
||||
if (nb + 2 < N_SLABS) {
|
||||
load_b_slab((nb % 2 == 0) ? sB0 : sB1, gB, p.b_s2, nb + 2, tid, policy_keep);
|
||||
cp_async_fence();
|
||||
}
|
||||
stage_flush(ss.c_stage.data(), gO, p.o_s0, nb * BN, m_residue, tid);
|
||||
} else {
|
||||
__syncthreads();
|
||||
if (nb + 2 < N_SLABS) {
|
||||
load_b_slab((nb % 2 == 0) ? sB0 : sB1, gB, p.b_s2, nb + 2, tid, policy_keep);
|
||||
cp_async_fence();
|
||||
}
|
||||
store_slab_direct(acc, gO, p.o_s0, nb * BN, row_base, col_base, m_residue);
|
||||
}
|
||||
if (nb + 1 < N_SLABS) {
|
||||
// Slab nb+1 resident (leave the nb+2 refill in flight, if any),
|
||||
// then commit the next round's gemm. The barrier also separates
|
||||
// this round's stage_flush reads from the next stage_write.
|
||||
if (nb + 2 < N_SLABS) {
|
||||
cp_async_wait<1>();
|
||||
} else {
|
||||
cp_async_wait<0>();
|
||||
}
|
||||
__syncthreads();
|
||||
gemm_ss(tiled_mma, sA, (nb % 2 == 0) ? sB1 : sB0, acc, tid);
|
||||
warpgroup_commit_batch();
|
||||
}
|
||||
};
|
||||
if constexpr (N_SLABS <= 4) {
|
||||
CUTE_UNROLL
|
||||
for (int nb = 0; nb < N_SLABS; ++nb) {
|
||||
round_body(nb);
|
||||
}
|
||||
} else {
|
||||
// Fully unrolling 8 rounds blows the 128-register budget (2 CTAs/SM
|
||||
// launch bound) and spills to local memory.
|
||||
CUTE_NO_UNROLL
|
||||
for (int nb = 0; nb < N_SLABS; ++nb) {
|
||||
round_body(nb);
|
||||
}
|
||||
}
|
||||
}
|
||||
#else
|
||||
if (cute::thread0()) {
|
||||
CUTE_INVALID_CONTROL_PATH("qprep_bf16_fp8_sm90 only supports sm90");
|
||||
}
|
||||
#endif
|
||||
}
|
||||
|
||||
// -------------------------------------------------------------------------
|
||||
// Host-side launch
|
||||
// -------------------------------------------------------------------------
|
||||
static void run(const QprepBf16Fp8Sm90Params& p) {
|
||||
QPREP_ASSERT(p.num_tokens > 0);
|
||||
QPREP_ASSERT(p.num_heads > 0);
|
||||
|
||||
auto kernel = &qprep_bf16_fp8_kernel<QprepBf16Fp8Kernel<K_DIM>>;
|
||||
constexpr size_t smem_size = sizeof(SharedStorage);
|
||||
static bool attr_set = [&]() {
|
||||
QPREP_CUDA_CHECK(cudaFuncSetAttribute(kernel, cudaFuncAttributeMaxDynamicSharedMemorySize, smem_size));
|
||||
return true;
|
||||
}();
|
||||
(void)attr_set;
|
||||
|
||||
dim3 grid(ceil_div_i(p.num_tokens, BM), p.num_heads, 1);
|
||||
kernel<<<grid, NUM_THREADS, smem_size, p.stream>>>(p);
|
||||
QPREP_CUDA_CHECK(cudaGetLastError());
|
||||
}
|
||||
};
|
||||
|
||||
template <typename Kernel>
|
||||
__global__ void __launch_bounds__(Kernel::NUM_THREADS, Kernel::MIN_CTAS)
|
||||
qprep_bf16_fp8_kernel(__grid_constant__ const QprepBf16Fp8Sm90Params params) {
|
||||
Kernel::devfunc(params);
|
||||
}
|
||||
|
||||
template <int K_DIM>
|
||||
void run_qprep_bf16_fp8_sm90(const QprepBf16Fp8Sm90Params& params) {
|
||||
QprepBf16Fp8Kernel<K_DIM>::run(params);
|
||||
}
|
||||
|
||||
} // namespace qprep_sm90
|
||||
@@ -0,0 +1,52 @@
|
||||
/* Copyright 2026 SGLang Team. All Rights Reserved.
|
||||
|
||||
Licensed under the Apache License, Version 2.0 (the "License");
|
||||
you may not use this file except in compliance with the License.
|
||||
You may obtain a copy of the License at
|
||||
|
||||
http://www.apache.org/licenses/LICENSE-2.0
|
||||
|
||||
Unless required by applicable law or agreed to in writing, software
|
||||
distributed under the License is distributed on an "AS IS" BASIS,
|
||||
WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
See the License for the specific language governing permissions and
|
||||
limitations under the License.
|
||||
==============================================================================*/
|
||||
|
||||
// Parameters for the SM90 Q8KV8 born-fp8 q-prep kernel (absorbed-q bmm +
|
||||
// nope/rope concat + fp32 -> bf16 -> fp8_e4m3 cast). All strides are in
|
||||
// ELEMENTS of the respective tensor's dtype (fp8 strides == byte strides).
|
||||
|
||||
#pragma once
|
||||
|
||||
#include <cstdint>
|
||||
#include <cuda_runtime.h>
|
||||
|
||||
struct QprepBf16Fp8Sm90Params {
|
||||
int num_tokens; // T (runtime; m-tiles are masked)
|
||||
int num_heads; // H (grid dim)
|
||||
|
||||
// q_nope: [T, H, K] bf16 (strided view OK; innermost dim contiguous)
|
||||
const void* q_nope;
|
||||
int64_t a_s0, a_s1;
|
||||
|
||||
// w_kc: [H, K, N] bf16 with K contiguous (stride(1) == 1; production layout
|
||||
// is (K*N, 1, K), i.e. the N-major absorbed weight)
|
||||
const void* w_kc;
|
||||
int64_t b_s0, b_s2;
|
||||
|
||||
// q_rope: [T, H, R] bf16 (strided view OK; innermost dim contiguous)
|
||||
const void* q_rope;
|
||||
int64_t r_s0, r_s1;
|
||||
// 16B-aligned rope rows (base pointer and both strides) -> uint4 loads
|
||||
bool rope_vec16;
|
||||
|
||||
// out: [T, pad_heads, N + R] fp8_e4m3; only [:, :H, :] is written
|
||||
void* out;
|
||||
int64_t o_s0, o_s1;
|
||||
// 16B-aligned out rows (base pointer and both strides) -> smem-staged
|
||||
// coalesced uint4 stores for the nope half (else direct u16 stores)
|
||||
bool out_vec16;
|
||||
|
||||
cudaStream_t stream;
|
||||
};
|
||||
@@ -199,4 +199,45 @@ void sparse_prefill_q8kv8_dispatch_full(
|
||||
_run_q8kv8(params, true, true);
|
||||
}
|
||||
|
||||
void sparse_prefill_q8kv8_dispatch_topk_length(
|
||||
tvm::ffi::TensorView q,
|
||||
tvm::ffi::TensorView kv,
|
||||
tvm::ffi::TensorView indices,
|
||||
tvm::ffi::TensorView q_scale,
|
||||
tvm::ffi::TensorView kv_scale,
|
||||
tvm::ffi::TensorView topk_length,
|
||||
tvm::ffi::TensorView out,
|
||||
tvm::ffi::TensorView max_logits,
|
||||
tvm::ffi::TensorView lse,
|
||||
int64_t s_q_val,
|
||||
int64_t s_kv_val,
|
||||
int64_t h_q_val,
|
||||
int64_t h_kv_val,
|
||||
int64_t d_qk_val,
|
||||
int64_t d_v_val,
|
||||
int64_t topk_val,
|
||||
double sm_scale_val,
|
||||
int64_t cuda_stream) {
|
||||
SparseMlaQ8Kv8PrefillParams params = _make_common_params(
|
||||
q,
|
||||
kv,
|
||||
indices,
|
||||
q_scale,
|
||||
kv_scale,
|
||||
out,
|
||||
max_logits,
|
||||
lse,
|
||||
s_q_val,
|
||||
s_kv_val,
|
||||
h_q_val,
|
||||
h_kv_val,
|
||||
d_qk_val,
|
||||
d_v_val,
|
||||
topk_val,
|
||||
sm_scale_val,
|
||||
cuda_stream);
|
||||
params.topk_length = static_cast<int*>(topk_length.data_ptr());
|
||||
_run_q8kv8(params, true, false);
|
||||
}
|
||||
|
||||
} // namespace
|
||||
|
||||
@@ -1,3 +1,5 @@
|
||||
from typing import Optional
|
||||
|
||||
import torch
|
||||
import triton
|
||||
import triton.language as tl
|
||||
@@ -285,11 +287,19 @@ def _dequantize_k_cache_paged_kernel(
|
||||
tl.store(dst_ptr, data, mask=mask)
|
||||
|
||||
|
||||
# Tokens handled by one program of the vectorized gather kernel. 4 tokens
|
||||
# x 512 fp8 nope elements = 2048 elements per program: with num_warps=4
|
||||
# (128 threads) that is 16 fp8 elements per thread, which Triton emits as
|
||||
# a single 16-byte vectorized load/store per thread.
|
||||
_GATHER_TOKENS_PER_PROG = 4
|
||||
|
||||
|
||||
def gather_dequant_requant_fp8_paged(
|
||||
quant_k_cache: torch.Tensor,
|
||||
page_table_1_flattened: torch.Tensor,
|
||||
group_size: int = 128,
|
||||
extra_rows: int = 0,
|
||||
out: Optional[torch.Tensor] = None,
|
||||
) -> torch.Tensor:
|
||||
"""Gather paged fp8 KV tokens and re-pack into flat [576] fp8 layout.
|
||||
|
||||
@@ -300,6 +310,13 @@ def gather_dequant_requant_fp8_paged(
|
||||
Rope is cast bf16->fp8. The whole operation is fused into a single
|
||||
Triton kernel to avoid allocating an intermediate bf16 buffer.
|
||||
|
||||
The kernel writes EVERY byte of rows [0, num_tokens) and zero-fills
|
||||
rows [num_tokens, num_tokens + extra_rows) (the -1-sentinel landing
|
||||
pad required by the SM90 sparse MLA Q8KV8 kernel, which clamps each
|
||||
-1 topk slot ``offs`` to distinct row ``num_tokens + offs``). The
|
||||
destination therefore needs no pre-zeroing, which allows passing a
|
||||
persistent (dirty) buffer via ``out``.
|
||||
|
||||
Args:
|
||||
quant_k_cache: [total_num_tokens, 1, 656] fp8_e4m3fn
|
||||
page_table_1_flattened: [num_tokens] int32
|
||||
@@ -308,6 +325,9 @@ def gather_dequant_requant_fp8_paged(
|
||||
the end of the output (used by the SM90 sparse MLA Q8KV8
|
||||
kernel which over-reads past end-of-buffer for masked
|
||||
indices)
|
||||
out: optional pre-allocated destination of shape
|
||||
[num_tokens + extra_rows, 1, 576] (or [.., 576]) fp8_e4m3fn,
|
||||
contiguous. Contents may be arbitrary (fully overwritten).
|
||||
Returns:
|
||||
output: [num_tokens + extra_rows, 1, 576] fp8_e4m3fn
|
||||
"""
|
||||
@@ -323,11 +343,162 @@ def gather_dequant_requant_fp8_paged(
|
||||
out_dim = dim_nope + dim_rope # 576
|
||||
assert num_tiles * group_size == dim_nope
|
||||
|
||||
total_rows = num_tokens + extra_rows
|
||||
if out is None:
|
||||
# No zero-fill needed: the kernel overwrites every byte of the
|
||||
# data rows and zero-fills the pad rows itself.
|
||||
output = torch.empty(
|
||||
(total_rows, 1, out_dim),
|
||||
dtype=torch.float8_e4m3fn,
|
||||
device=quant_k_cache.device,
|
||||
)
|
||||
else:
|
||||
assert out.dtype == torch.float8_e4m3fn
|
||||
assert out.device == quant_k_cache.device
|
||||
assert out.is_contiguous()
|
||||
assert out.numel() == total_rows * out_dim, (
|
||||
f"out buffer has {out.numel()} elements, expected "
|
||||
f"{total_rows} x {out_dim} = {total_rows * out_dim}"
|
||||
)
|
||||
output = out.view(total_rows, 1, out_dim)
|
||||
|
||||
if total_rows == 0:
|
||||
return output
|
||||
|
||||
input_nope_q = quant_k_cache[:, :dim_nope]
|
||||
input_nope_s = quant_k_cache[:, dim_nope : dim_nope + num_tiles * 4].view(
|
||||
torch.float32
|
||||
)
|
||||
input_rope = quant_k_cache[:, dim_nope + num_tiles * 4 :].view(torch.bfloat16)
|
||||
|
||||
grid = (triton.cdiv(total_rows, _GATHER_TOKENS_PER_PROG),)
|
||||
_gather_dequant_requant_fp8_paged_vec_kernel[grid](
|
||||
output,
|
||||
input_nope_q,
|
||||
input_nope_s,
|
||||
input_rope,
|
||||
page_table_1_flattened,
|
||||
num_tokens,
|
||||
total_rows,
|
||||
output.stride(0),
|
||||
input_nope_q.stride(0),
|
||||
input_nope_s.stride(0),
|
||||
input_rope.stride(0),
|
||||
NUM_NOPE_BLOCKS=num_tiles,
|
||||
GROUP_SIZE=group_size,
|
||||
DIM_NOPE=dim_nope,
|
||||
DIM_ROPE=dim_rope,
|
||||
TOKENS_PER_PROG=_GATHER_TOKENS_PER_PROG,
|
||||
num_warps=4,
|
||||
)
|
||||
|
||||
return output
|
||||
|
||||
|
||||
@triton.jit
|
||||
def _gather_dequant_requant_fp8_paged_vec_kernel(
|
||||
output_ptr,
|
||||
input_nope_q_ptr,
|
||||
input_nope_s_ptr,
|
||||
input_rope_ptr,
|
||||
page_table_1_ptr,
|
||||
num_tokens: int,
|
||||
total_rows: int,
|
||||
output_stride_0: int,
|
||||
input_nope_q_stride_0: int,
|
||||
input_nope_s_stride_0: int,
|
||||
input_rope_stride_0: int,
|
||||
NUM_NOPE_BLOCKS: tl.constexpr,
|
||||
GROUP_SIZE: tl.constexpr,
|
||||
DIM_NOPE: tl.constexpr,
|
||||
DIM_ROPE: tl.constexpr,
|
||||
TOKENS_PER_PROG: tl.constexpr,
|
||||
):
|
||||
"""Vectorized fused gather + dequant(per-group) + requant(per-tensor).
|
||||
|
||||
One program handles TOKENS_PER_PROG consecutive output rows (full
|
||||
576-byte rows each), instead of the legacy one-program-per-(token,
|
||||
128-elem-slice) layout, so each thread moves 16 contiguous fp8 bytes
|
||||
per load/store. Rows >= num_tokens (the -1-sentinel landing pad) are
|
||||
zero-filled without touching the KV cache. Per-element math is
|
||||
bit-identical to the legacy kernel: fp8 -> f32, * f32 group scale,
|
||||
-> fp8 (nope); bf16 -> fp8 (rope).
|
||||
"""
|
||||
pid = tl.program_id(0)
|
||||
offs_t = pid * TOKENS_PER_PROG + tl.arange(0, TOKENS_PER_PROG) # [T]
|
||||
row_in_range = offs_t < total_rows
|
||||
is_real = offs_t < num_tokens
|
||||
# Masked lanes (pad rows) never touch memory; `other=0` keeps the
|
||||
# address arithmetic in-bounds-irrelevant.
|
||||
paged = tl.load(page_table_1_ptr + offs_t, mask=is_real, other=0).to(tl.int64)
|
||||
# 64-bit output row offsets: total_rows * 576 can exceed int32 for
|
||||
# very large gathered buffers.
|
||||
offs_t64 = offs_t.to(tl.int64)
|
||||
|
||||
offs_g = tl.arange(0, NUM_NOPE_BLOCKS) # [G] dequant groups
|
||||
offs_i = tl.arange(0, GROUP_SIZE) # [I] elems within a group
|
||||
|
||||
# a. nope: [T, G, I] fp8 block; the (G, I) plane spans the contiguous
|
||||
# DIM_NOPE bytes of one cache row.
|
||||
ptr_q = (
|
||||
input_nope_q_ptr
|
||||
+ paged[:, None, None] * input_nope_q_stride_0
|
||||
+ offs_g[None, :, None] * GROUP_SIZE
|
||||
+ offs_i[None, None, :]
|
||||
)
|
||||
y_q = tl.load(ptr_q, mask=is_real[:, None, None], other=0.0).to(tl.float32)
|
||||
ptr_s = input_nope_s_ptr + paged[:, None] * input_nope_s_stride_0 + offs_g[None, :]
|
||||
y_s = tl.load(ptr_s, mask=is_real[:, None], other=0.0)
|
||||
# dequant -> f32 -> requant to fp8; pad rows: (0 * 0) -> +0 -> byte 0x00
|
||||
y = (y_q * y_s[:, :, None]).to(tl.float8e4nv)
|
||||
dst_q = (
|
||||
output_ptr
|
||||
+ offs_t64[:, None, None] * output_stride_0
|
||||
+ offs_g[None, :, None] * GROUP_SIZE
|
||||
+ offs_i[None, None, :]
|
||||
)
|
||||
tl.store(dst_q, y, mask=row_in_range[:, None, None])
|
||||
|
||||
# b. rope: [T, R] bf16 -> fp8; pad rows: 0.0 -> byte 0x00
|
||||
offs_r = tl.arange(0, DIM_ROPE)
|
||||
src_r = input_rope_ptr + paged[:, None] * input_rope_stride_0 + offs_r[None, :]
|
||||
data = tl.load(src_r, mask=is_real[:, None], other=0.0).to(tl.float8e4nv)
|
||||
dst_r = (
|
||||
output_ptr + offs_t64[:, None] * output_stride_0 + DIM_NOPE + offs_r[None, :]
|
||||
)
|
||||
tl.store(dst_r, data, mask=row_in_range[:, None])
|
||||
|
||||
|
||||
def gather_dequant_requant_fp8_paged_legacy(
|
||||
quant_k_cache: torch.Tensor,
|
||||
page_table_1_flattened: torch.Tensor,
|
||||
group_size: int = 128,
|
||||
extra_rows: int = 0,
|
||||
) -> torch.Tensor:
|
||||
"""Legacy (pre-vectorization) gather + dequant + requant.
|
||||
|
||||
Kept as the bit-exactness / performance reference for
|
||||
``gather_dequant_requant_fp8_paged`` (see
|
||||
``benchmark/kernels/deepseek/benchmark_q8kv8_kv_gather.py``). Allocates and zero-fills the
|
||||
full destination each call, then launches one program per
|
||||
(token, 128-elem slice).
|
||||
"""
|
||||
dim_quant = quant_k_cache.shape[-1]
|
||||
assert dim_quant == 656
|
||||
quant_k_cache = quant_k_cache.view((-1, dim_quant))
|
||||
|
||||
num_tokens = page_table_1_flattened.shape[0]
|
||||
assert quant_k_cache.dtype == torch.float8_e4m3fn
|
||||
dim_nope = 512
|
||||
dim_rope = 64
|
||||
num_tiles = dim_nope // group_size # 4
|
||||
out_dim = dim_nope + dim_rope # 576
|
||||
assert num_tiles * group_size == dim_nope
|
||||
|
||||
total_rows = num_tokens + extra_rows
|
||||
# Allocate a fresh zero-filled buffer. The extra landing-pad rows at
|
||||
# the tail must read as zeros (the kernel may over-read past
|
||||
# num_tokens for masked indices). A future optimization could cache
|
||||
# this buffer but baseline allocates fresh.
|
||||
# num_tokens for masked indices).
|
||||
output = torch.zeros(
|
||||
(total_rows, 1, out_dim),
|
||||
dtype=torch.float8_e4m3fn,
|
||||
@@ -419,3 +590,67 @@ def _gather_dequant_requant_fp8_paged_kernel(
|
||||
|
||||
if __name__ == "__main__":
|
||||
raise Exception("UT is in quant_k_cache.py")
|
||||
|
||||
|
||||
@triton.jit
|
||||
def _concat_cast_kv_fp8_pad_kernel(
|
||||
out_ptr,
|
||||
k_ptr,
|
||||
kr_ptr,
|
||||
num_tokens,
|
||||
k_stride,
|
||||
kr_stride,
|
||||
NOPE: tl.constexpr,
|
||||
ROPE: tl.constexpr,
|
||||
):
|
||||
"""Row program: real rows write cast(k)||cast(k_rope); pad-band rows
|
||||
write zeros (the -1-sentinel landing pad the kernel's clamp maps to)."""
|
||||
row = tl.program_id(0).to(tl.int64)
|
||||
offs_n = tl.arange(0, NOPE)
|
||||
offs_r = tl.arange(0, ROPE)
|
||||
head = NOPE + ROPE
|
||||
if row < num_tokens:
|
||||
v_n = tl.load(k_ptr + row * k_stride + offs_n)
|
||||
tl.store(out_ptr + row * head + offs_n, v_n.to(tl.float8e4nv))
|
||||
v_r = tl.load(kr_ptr + row * kr_stride + offs_r)
|
||||
tl.store(out_ptr + row * head + NOPE + offs_r, v_r.to(tl.float8e4nv))
|
||||
else:
|
||||
zero_n = tl.zeros([NOPE], dtype=tl.float32).to(tl.float8e4nv)
|
||||
zero_r = tl.zeros([ROPE], dtype=tl.float32).to(tl.float8e4nv)
|
||||
tl.store(out_ptr + row * head + offs_n, zero_n)
|
||||
tl.store(out_ptr + row * head + NOPE + offs_r, zero_r)
|
||||
|
||||
|
||||
def concat_cast_kv_fp8_pad(
|
||||
out: torch.Tensor,
|
||||
k: torch.Tensor,
|
||||
k_rope: torch.Tensor,
|
||||
num_tokens: int,
|
||||
) -> torch.Tensor:
|
||||
"""Fused non-prefix Q8KV8 KV prep: cast-concat k (nope latent) and k_rope
|
||||
directly into the persistent fp8 kv_buf and zero the trailing pad band —
|
||||
replaces the bf16 `_cat` materialization + `.copy_` cast + `.zero_()`
|
||||
tail (3 kernels + one [tokens, 576] bf16 alloc). Same bf16->fp8
|
||||
store-cast the gather kernel uses (bit-identical bytes).
|
||||
|
||||
``out``: [total_rows, 576] fp8 slice (total_rows = num_tokens + pad band);
|
||||
``k``: [num_tokens, NOPE] bf16 view; ``k_rope``: [num_tokens, ROPE] bf16.
|
||||
"""
|
||||
total_rows, head = out.shape
|
||||
nope = k.shape[-1]
|
||||
rope = k_rope.shape[-1]
|
||||
assert head == nope + rope and out.dtype == torch.float8_e4m3fn
|
||||
k2 = k.view(num_tokens, nope)
|
||||
kr2 = k_rope.view(num_tokens, rope)
|
||||
assert k2.stride(-1) == 1 and kr2.stride(-1) == 1
|
||||
_concat_cast_kv_fp8_pad_kernel[(total_rows,)](
|
||||
out,
|
||||
k2,
|
||||
kr2,
|
||||
num_tokens,
|
||||
k2.stride(0),
|
||||
kr2.stride(0),
|
||||
NOPE=nope,
|
||||
ROPE=rope,
|
||||
)
|
||||
return out
|
||||
|
||||
@@ -7,6 +7,36 @@ from collections.abc import Callable
|
||||
import torch
|
||||
|
||||
|
||||
def _restore_row_stride(logits: torch.Tensor) -> torch.Tensor:
|
||||
"""Undo PyTorch's DLPack stride normalization on single-row DeepGEMM logits.
|
||||
|
||||
DeepGEMM returns paged-MQA logits as a row-padded view:
|
||||
``torch.empty(num_rows, aligned_len)[:, :max_len]`` with ``aligned_len``
|
||||
256-element (1024-byte) aligned. tvm-ffi builds of DeepGEMM (sgl-deep-gemm
|
||||
>= 0.1.x) round-trip that view through DLPack on return, and PyTorch's
|
||||
DLPack *exporter* rewrites the stride of every size<2 dim to 1 whenever it
|
||||
differs from the packed expectation (pytorch/pytorch#83158). A one-row
|
||||
result whose row is actually padded (``max_len % 256 != 0``, e.g. any
|
||||
``model context_len + 4`` page-table width at bs=1 decode capture) therefore
|
||||
arrives with ``stride() == (1, 1)`` instead of ``(aligned_len, 1)``, which
|
||||
violates the fused top-k v2 kernel ABI (``score_stride % 4 == 0``, enforced
|
||||
both in ``dsa_topk_backend._topk_transform_v2_paged`` and by the kernel's
|
||||
own RuntimeCheck).
|
||||
|
||||
For ``num_rows <= 1`` the row stride is semantically arbitrary (row 0 is
|
||||
the only row ever addressed and ``(num_rows - 1) * stride(0)`` contributes
|
||||
nothing to the storage extent), so restoring a 16-byte-aligned value is a
|
||||
pure metadata rewrite: same storage, same data pointer, no copy and no
|
||||
kernel launch -- trivially CUDA-graph-capture-safe. Multi-row results keep
|
||||
their true strides through DLPack (no size<2 dim) and pass through
|
||||
untouched.
|
||||
"""
|
||||
if logits.shape[0] <= 1 and logits.stride(0) % 4 != 0:
|
||||
width = logits.shape[1]
|
||||
logits = logits.as_strided((logits.shape[0], width), ((width + 3) // 4 * 4, 1))
|
||||
return logits
|
||||
|
||||
|
||||
def deepgemm_paged_mqa_logits_native(
|
||||
fp8_paged_mqa_logits_fn: Callable[..., torch.Tensor],
|
||||
q_fp8: torch.Tensor,
|
||||
@@ -23,7 +53,7 @@ def deepgemm_paged_mqa_logits_native(
|
||||
) -> torch.Tensor:
|
||||
# block_tables[::next_n] de-expands the caller's repeat_interleave without a
|
||||
# copy (DeepGEMM only checks `stride(1) == 1`).
|
||||
return fp8_paged_mqa_logits_fn(
|
||||
logits = fp8_paged_mqa_logits_fn(
|
||||
q_fp8[:q_offset].view(B, next_n, q_fp8.shape[1], q_fp8.shape[2]),
|
||||
kv_cache_fp8,
|
||||
weights[:q_offset],
|
||||
@@ -33,6 +63,7 @@ def deepgemm_paged_mqa_logits_native(
|
||||
max_seq_len,
|
||||
clean_logits=False,
|
||||
)
|
||||
return _restore_row_stride(logits)
|
||||
|
||||
|
||||
def deepgemm_paged_mqa_logits_split(
|
||||
@@ -48,7 +79,7 @@ def deepgemm_paged_mqa_logits_split(
|
||||
q_offset: int,
|
||||
) -> torch.Tensor:
|
||||
q_fp8 = q_fp8.unsqueeze(1)
|
||||
return fp8_paged_mqa_logits_fn(
|
||||
logits = fp8_paged_mqa_logits_fn(
|
||||
q_fp8[:q_offset],
|
||||
kv_cache_fp8,
|
||||
weights[:q_offset],
|
||||
@@ -58,6 +89,7 @@ def deepgemm_paged_mqa_logits_split(
|
||||
max_seq_len,
|
||||
clean_logits=False,
|
||||
)
|
||||
return _restore_row_stride(logits)
|
||||
|
||||
|
||||
def aiter_paged_mqa_logits(
|
||||
|
||||
@@ -0,0 +1,134 @@
|
||||
"""JIT-compiled SM90 (Hopper) kernel for the Q8KV8 born-fp8 q-prep.
|
||||
|
||||
Fuses the per-head absorbed-q bmm (q_nope [T, H, K] bf16 x w_kc [H, K, N]
|
||||
bf16, fp32 accumulate), the nope/rope concat, and the bf16 -> fp8_e4m3 cast
|
||||
into one hand-written WGMMA kernel. CUDA replacement for the Triton
|
||||
``absorbed_bmm_concat_cast_q_fp8`` (triton_ops/cache_ops.py) with the
|
||||
identical fp32 -> bf16 -> fp8 epilogue rounding chain; the rope half is
|
||||
bit-exact vs ``concat_and_cast_q_fp8_pad``.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import TYPE_CHECKING
|
||||
|
||||
import torch
|
||||
|
||||
from sglang.kernel_api_logging import debug_kernel_api
|
||||
from sglang.kernels.jit.utils import cache_once, load_jit, override_jit_cuda_arch
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from tvm_ffi.module import Module
|
||||
|
||||
|
||||
N_LORA = 512 # kv_lora_rank (nope output dim)
|
||||
ROPE_DIM = 64 # qk_rope_head_dim
|
||||
|
||||
|
||||
@cache_once
|
||||
def _jit_qprep_bf16_fp8_module() -> Module:
|
||||
if torch.cuda.get_device_capability()[0] != 9:
|
||||
raise RuntimeError("qprep_bf16_fp8_sm90 requires an SM90 (Hopper) GPU")
|
||||
with override_jit_cuda_arch(9, 0, "a"):
|
||||
return load_jit(
|
||||
"qprep_bf16_fp8_sm90",
|
||||
cuda_files=["qprep_bf16_fp8_sm90/entry.cuh"],
|
||||
cuda_wrappers=[("dispatch", "qprep_bf16_fp8_dispatch")],
|
||||
# Same minimal flag set as the sparse_mla_q8kv8_prefill_sm90 JIT
|
||||
# build (per-flag ablation there showed the rest are no-ops).
|
||||
extra_cuda_cflags=[
|
||||
"-O3",
|
||||
"-DNDEBUG",
|
||||
"-DCUTE_USE_PACKED_TUPLE=1",
|
||||
"-DCUTLASS_ENABLE_TENSOR_CORE_MMA=1",
|
||||
"--use_fast_math",
|
||||
],
|
||||
extra_dependencies=["cutlass"],
|
||||
)
|
||||
|
||||
|
||||
# torch._C._cuda_getCurrentRawStream returns the cudaStream_t pointer expected
|
||||
# by the JIT wrapper (see sparse_mla_q8kv8_prefill_sm90.py).
|
||||
_get_current_stream_raw = torch._C._cuda_getCurrentRawStream
|
||||
|
||||
|
||||
@debug_kernel_api
|
||||
def q8kv8_qprep_fwd(
|
||||
q_fp8_pad: torch.Tensor,
|
||||
q_nope: torch.Tensor,
|
||||
w_kc: torch.Tensor,
|
||||
q_rope: torch.Tensor,
|
||||
num_heads: int,
|
||||
) -> None:
|
||||
"""Fused absorbed-q bmm + nope/rope concat + bf16->fp8 cast ("born fp8" q).
|
||||
|
||||
Mirrors the contract of ``absorbed_bmm_concat_cast_q_fp8``:
|
||||
|
||||
* ``q_fp8_pad``: [num_tokens, pad_heads, N + ROPE] fp8_e4m3 destination;
|
||||
only ``[:, :num_heads, :]`` is written.
|
||||
* ``q_nope``: [num_tokens, H, K] bf16 pre-absorb q (strided views OK).
|
||||
* ``w_kc``: [H, K, N] bf16 absorbed weight with K contiguous
|
||||
(``stride(1) == 1``, the production N-major layout).
|
||||
* ``q_rope``: [num_tokens, H, ROPE] bf16 post-rope q (strided views OK).
|
||||
|
||||
K (``qk_nope_head_dim``) must be 128 or 192. Extra restrictions vs the
|
||||
Triton kernel (all satisfied by the production layouts): 16-byte aligned
|
||||
q_nope/w_kc base pointers, q_nope/w_kc strides that are multiples of 8
|
||||
elements, and even q_fp8_pad row/head strides.
|
||||
"""
|
||||
num_tokens, _, k_dim = q_nope.shape
|
||||
n_dim = w_kc.shape[-1]
|
||||
rope_dim = q_rope.shape[-1]
|
||||
assert q_fp8_pad.dtype == torch.float8_e4m3fn
|
||||
assert q_nope.dtype == torch.bfloat16 and w_kc.dtype == torch.bfloat16
|
||||
assert q_rope.dtype == torch.bfloat16
|
||||
assert q_nope.is_cuda and w_kc.is_cuda and q_rope.is_cuda and q_fp8_pad.is_cuda
|
||||
assert q_nope.shape[1] == num_heads and q_rope.shape[1] == num_heads
|
||||
assert w_kc.shape[0] == num_heads and w_kc.shape[1] == k_dim
|
||||
assert q_fp8_pad.shape[0] >= num_tokens and q_fp8_pad.shape[1] >= num_heads
|
||||
assert q_fp8_pad.shape[2] == n_dim + rope_dim
|
||||
assert k_dim in (128, 192), "CUDA q-prep supports K in {128, 192}"
|
||||
assert n_dim == N_LORA and rope_dim == ROPE_DIM
|
||||
# Innermost-contiguous requirements (same as the Triton kernel).
|
||||
assert q_nope.stride(2) == 1 and q_rope.stride(2) == 1
|
||||
assert q_fp8_pad.stride(2) == 1
|
||||
# CUDA-kernel-specific layout requirements (production layouts satisfy
|
||||
# all of these; the Triton kernel stays the general-strides fallback).
|
||||
assert w_kc.stride(1) == 1, "w_kc must have K contiguous (N-major layout)"
|
||||
assert q_nope.data_ptr() % 16 == 0 and w_kc.data_ptr() % 16 == 0
|
||||
assert q_nope.stride(0) % 8 == 0 and q_nope.stride(1) % 8 == 0
|
||||
assert w_kc.stride(0) % 8 == 0 and w_kc.stride(2) % 8 == 0
|
||||
assert q_fp8_pad.stride(0) % 2 == 0 and q_fp8_pad.stride(1) % 2 == 0
|
||||
|
||||
rope_vec16 = (
|
||||
q_rope.data_ptr() % 16 == 0
|
||||
and q_rope.stride(0) % 8 == 0
|
||||
and q_rope.stride(1) % 8 == 0
|
||||
)
|
||||
out_vec16 = (
|
||||
q_fp8_pad.data_ptr() % 16 == 0
|
||||
and q_fp8_pad.stride(0) % 16 == 0
|
||||
and q_fp8_pad.stride(1) % 16 == 0
|
||||
)
|
||||
|
||||
module = _jit_qprep_bf16_fp8_module()
|
||||
module.dispatch(
|
||||
q_nope,
|
||||
w_kc,
|
||||
q_rope,
|
||||
q_fp8_pad,
|
||||
num_tokens,
|
||||
num_heads,
|
||||
k_dim,
|
||||
q_nope.stride(0),
|
||||
q_nope.stride(1),
|
||||
w_kc.stride(0),
|
||||
w_kc.stride(2),
|
||||
q_rope.stride(0),
|
||||
q_rope.stride(1),
|
||||
q_fp8_pad.stride(0),
|
||||
q_fp8_pad.stride(1),
|
||||
int(rope_vec16),
|
||||
int(out_vec16),
|
||||
_get_current_stream_raw(q_nope.device.index),
|
||||
)
|
||||
@@ -68,6 +68,7 @@ def _jit_sparse_mla_q8kv8_prefill_module() -> Module:
|
||||
cuda_wrappers=[
|
||||
("dispatch", "sparse_prefill_q8kv8_dispatch"),
|
||||
("dispatch_full", "sparse_prefill_q8kv8_dispatch_full"),
|
||||
("dispatch_topk_length", "sparse_prefill_q8kv8_dispatch_topk_length"),
|
||||
],
|
||||
extra_cuda_cflags=_q8kv8_cuda_flags(),
|
||||
extra_dependencies=["cutlass"],
|
||||
@@ -86,6 +87,7 @@ def _get_entries() -> tuple:
|
||||
_resolved_entries = (
|
||||
m["dispatch"],
|
||||
m["dispatch_full"],
|
||||
m["dispatch_topk_length"],
|
||||
)
|
||||
return _resolved_entries
|
||||
|
||||
@@ -146,7 +148,7 @@ def _sparse_mla_q8kv8_prefill_op(
|
||||
sm_scale: float,
|
||||
cuda_stream: int,
|
||||
) -> None:
|
||||
dispatch_fn, _ = _get_entries()
|
||||
dispatch_fn, _, _ = _get_entries()
|
||||
dispatch_fn(
|
||||
q,
|
||||
kv,
|
||||
@@ -193,7 +195,7 @@ def _sparse_mla_q8kv8_prefill_full_op(
|
||||
sm_scale: float,
|
||||
cuda_stream: int,
|
||||
) -> None:
|
||||
_, dispatch_full_fn = _get_entries()
|
||||
_, dispatch_full_fn, _ = _get_entries()
|
||||
dispatch_full_fn(
|
||||
q,
|
||||
kv,
|
||||
@@ -217,6 +219,53 @@ def _sparse_mla_q8kv8_prefill_full_op(
|
||||
)
|
||||
|
||||
|
||||
@register_custom_op(
|
||||
op_name="sparse_mla_q8kv8_prefill_topk_length",
|
||||
mutates_args=["out", "max_logits", "lse"],
|
||||
)
|
||||
def _sparse_mla_q8kv8_prefill_topk_length_op(
|
||||
q: torch.Tensor,
|
||||
kv: torch.Tensor,
|
||||
indices: torch.Tensor,
|
||||
q_scale: torch.Tensor,
|
||||
kv_scale: torch.Tensor,
|
||||
topk_length: torch.Tensor,
|
||||
out: torch.Tensor,
|
||||
max_logits: torch.Tensor,
|
||||
lse: torch.Tensor,
|
||||
s_q: int,
|
||||
s_kv: int,
|
||||
h_q: int,
|
||||
h_kv: int,
|
||||
d_qk: int,
|
||||
d_v: int,
|
||||
topk: int,
|
||||
sm_scale: float,
|
||||
cuda_stream: int,
|
||||
) -> None:
|
||||
_, _, dispatch_topk_length_fn = _get_entries()
|
||||
dispatch_topk_length_fn(
|
||||
q,
|
||||
kv,
|
||||
indices,
|
||||
q_scale,
|
||||
kv_scale,
|
||||
topk_length,
|
||||
out,
|
||||
max_logits,
|
||||
lse,
|
||||
s_q,
|
||||
s_kv,
|
||||
h_q,
|
||||
h_kv,
|
||||
d_qk,
|
||||
d_v,
|
||||
topk,
|
||||
sm_scale,
|
||||
cuda_stream,
|
||||
)
|
||||
|
||||
|
||||
@debug_kernel_api
|
||||
def sparse_mla_q8kv8_prefill_fwd(
|
||||
q: torch.Tensor, # [s_q, h_q, d_qk], float8_e4m3fn
|
||||
@@ -256,8 +305,8 @@ def sparse_mla_q8kv8_prefill_fwd(
|
||||
f"sparse_mla_q8kv8_prefill_fwd only supports d_v=512, got {d_v}"
|
||||
)
|
||||
|
||||
if (attn_sink is None) != (topk_length is None):
|
||||
raise ValueError("attn_sink and topk_length must be provided together")
|
||||
if attn_sink is not None and topk_length is None:
|
||||
raise ValueError("attn_sink requires topk_length to be provided as well")
|
||||
|
||||
device = q.device
|
||||
if out is None:
|
||||
@@ -305,6 +354,27 @@ def sparse_mla_q8kv8_prefill_fwd(
|
||||
sm_scale,
|
||||
cuda_stream,
|
||||
)
|
||||
elif topk_length is not None:
|
||||
_sparse_mla_q8kv8_prefill_topk_length_op(
|
||||
q,
|
||||
kv,
|
||||
indices,
|
||||
q_scale,
|
||||
kv_scale,
|
||||
topk_length,
|
||||
out,
|
||||
max_logits,
|
||||
lse,
|
||||
s_q,
|
||||
s_kv,
|
||||
h_q,
|
||||
h_kv,
|
||||
d_qk,
|
||||
d_v,
|
||||
topk,
|
||||
sm_scale,
|
||||
cuda_stream,
|
||||
)
|
||||
else:
|
||||
_sparse_mla_q8kv8_prefill_op(
|
||||
q,
|
||||
|
||||
@@ -30,6 +30,9 @@ from sglang.kernels.ops.kvcache.cache_ops import (
|
||||
from sglang.kernels.ops.kvcache.cache_ops import (
|
||||
launch_reshape_and_cache_flash as launch_reshape_and_cache_flash,
|
||||
)
|
||||
from sglang.kernels.ops.kvcache.cache_ops import (
|
||||
q8kv8_topk_length_from_indices as q8kv8_topk_length_from_indices,
|
||||
)
|
||||
from sglang.kernels.ops.kvcache.cache_ops import (
|
||||
reshape_and_cache_flash as reshape_and_cache_flash,
|
||||
)
|
||||
|
||||
@@ -322,6 +322,396 @@ def concat_and_cast_q_fp8_pad(q_fp8_pad, q_nope, q_rope, num_heads):
|
||||
)
|
||||
|
||||
|
||||
@triton.jit
|
||||
def absorbed_bmm_concat_cast_q_fp8_kernel(
|
||||
qout_ptr, # [num_tokens, pad_heads, N+ROPE] fp8 (dst; only [:, :H, :] written)
|
||||
a_ptr, # q_nope (pre-absorb) [num_tokens, H, K] bf16
|
||||
b_ptr, # w_kc [H, K, N] bf16 (any strides; typically N-major)
|
||||
rope_ptr, # q_rope (post-rope) [num_tokens, H, ROPE] bf16
|
||||
T, # num_tokens (runtime; masked)
|
||||
qout_s0,
|
||||
qout_s1,
|
||||
a_s0,
|
||||
a_s1,
|
||||
b_s0,
|
||||
b_s1,
|
||||
b_s2,
|
||||
rope_s0,
|
||||
rope_s1,
|
||||
K: tl.constexpr,
|
||||
N: tl.constexpr,
|
||||
ROPE: tl.constexpr,
|
||||
BLOCK_M: tl.constexpr,
|
||||
BLOCK_N: tl.constexpr,
|
||||
BLOCK_K: tl.constexpr,
|
||||
K_MODE: tl.constexpr,
|
||||
):
|
||||
# One program per (token-block, head): q_out[m, h, :N] = fp8(bf16(fp32(
|
||||
# q_nope[m, h, :K] @ w_kc[h, :K, :N]))) and q_out[m, h, N:] =
|
||||
# fp8(q_rope[m, h, :ROPE]). This makes q "born fp8": the absorbed bmm,
|
||||
# the nope/rope concat, and the bf16->fp8 cast collapse into one kernel,
|
||||
# so neither the bf16 q_nope_out ([H, T, N], written by cublas and re-read
|
||||
# by the concat-cast) nor the standalone concat-cast launch exist anymore.
|
||||
#
|
||||
# K handling (K_MODE selects the codegen for the nope-gemm K dimension;
|
||||
# every mode keeps the same fp32-accumulator -> bf16 -> fp8 epilogue):
|
||||
# 0 "single": BLOCK_K == K. Preload the whole [BLOCK_M, K] a-tile once,
|
||||
# one tl.dot per N-block — identical codegen to the original
|
||||
# power-of-2-only kernel (DeepSeek K=128). For non-power-of-2 K this
|
||||
# only compiles if the Triton build allows non-power-of-2 tl.arange
|
||||
# (Triton <= 3.5.x does NOT: "arange's range must be a power of 2").
|
||||
# 1 "loop": split-K loop, K % BLOCK_K == 0, BLOCK_K power of 2 >= 16
|
||||
# (e.g. K=192 with BLOCK_K=64 -> 3 iterations). The a-tile is
|
||||
# re-loaded per (N-block, K-block); slices are L1/L2-resident after
|
||||
# the first N-block, but the load/dot interleave costs bandwidth
|
||||
# (measured ~1372 GB/s vs ~2x that for mode 0 at K=128).
|
||||
# 2 "two_dot": K = BLOCK_K + (K - BLOCK_K), both power-of-2 halves
|
||||
# (192 = 128 + 64). Both a-tiles preload once before the N-loop;
|
||||
# each N-block issues two chained tl.dot into one fp32 accumulator.
|
||||
# No K-loop, no a re-reads — the direct generalization of mode 0.
|
||||
# 3 "three_dot": K = 3 * BLOCK_K (192 = 3 x 64). Same as mode 2 with
|
||||
# three preloaded a-tiles / three chained tl.dot per N-block; the
|
||||
# hoisted-loads analogue of mode 1 (identical fp32 add order).
|
||||
# 4 "pad": BLOCK_K = next_pow2(K) > K, k-masked loads (zero fill).
|
||||
# Single tl.dot per N-block; the padded zeros are exact fp32
|
||||
# additive identities so the result matches a K-wide single dot,
|
||||
# at the cost of BLOCK_K/K (e.g. 256/192 = 1.33x) extra MMA work.
|
||||
#
|
||||
# Rounding contract: the fp32 accumulator is rounded to bf16 first (the
|
||||
# same output rounding stage as the cublas bf16 bmm) and then converted
|
||||
# bf16->fp8 by the same implicit-store conversion the fused concat-cast
|
||||
# kernel uses. The split-K accumulator stays fp32 across all K-blocks,
|
||||
# so the rounding stages are identical in both layouts. The rope half is
|
||||
# a bit-exact copy of that kernel (loads the post-rope bf16, converts on
|
||||
# store). The nope half is NOT guaranteed bit-exact vs the default path:
|
||||
# tl.dot accumulates fp32 in a different order than cublas, so last-ulp
|
||||
# fp32 differences can occasionally flip the bf16 (and hence fp8)
|
||||
# rounding.
|
||||
pid_m = tl.program_id(0)
|
||||
h = tl.program_id(1)
|
||||
offs_m = pid_m * BLOCK_M + tl.arange(0, BLOCK_M)
|
||||
m_mask = offs_m < T
|
||||
# token-row offsets in int64: T * row-stride can exceed int32 (e.g. 128
|
||||
# heads x 576 dims x tens of thousands of tokens).
|
||||
offs_m64 = offs_m.to(tl.int64)
|
||||
qout_head = qout_ptr + offs_m64[:, None] * qout_s0 + h * qout_s1
|
||||
a_row = a_ptr + offs_m64[:, None] * a_s0 + h * a_s1
|
||||
b_head = b_ptr + h * b_s0
|
||||
if K_MODE == 0:
|
||||
# single-dot path (original kernel): BLOCK_K == K, preload a once.
|
||||
offs_k = tl.arange(0, BLOCK_K)
|
||||
a = tl.load(a_row + offs_k[None, :], mask=m_mask[:, None], other=0.0)
|
||||
for nb in tl.static_range(N // BLOCK_N):
|
||||
offs_n = nb * BLOCK_N + tl.arange(0, BLOCK_N)
|
||||
b = tl.load(b_head + offs_k[:, None] * b_s1 + offs_n[None, :] * b_s2)
|
||||
acc = tl.dot(a, b) # fp32 accumulator
|
||||
val = acc.to(tl.bfloat16) # cublas-equivalent bf16 output rounding
|
||||
# implicit bf16 -> fp8 conversion on store (same as the concat-cast)
|
||||
tl.store(qout_head + offs_n[None, :], val, mask=m_mask[:, None])
|
||||
elif K_MODE == 2:
|
||||
# two-dot preload: K split as BLOCK_K + (K - BLOCK_K), no K-loop.
|
||||
offs_k0 = tl.arange(0, BLOCK_K)
|
||||
offs_k1 = BLOCK_K + tl.arange(0, K - BLOCK_K)
|
||||
a0 = tl.load(a_row + offs_k0[None, :], mask=m_mask[:, None], other=0.0)
|
||||
a1 = tl.load(a_row + offs_k1[None, :], mask=m_mask[:, None], other=0.0)
|
||||
for nb in tl.static_range(N // BLOCK_N):
|
||||
offs_n = nb * BLOCK_N + tl.arange(0, BLOCK_N)
|
||||
b0 = tl.load(b_head + offs_k0[:, None] * b_s1 + offs_n[None, :] * b_s2)
|
||||
b1 = tl.load(b_head + offs_k1[:, None] * b_s1 + offs_n[None, :] * b_s2)
|
||||
acc = tl.dot(a0, b0) # fp32 accumulator
|
||||
acc = tl.dot(a1, b1, acc) # chained: stays fp32 across both dots
|
||||
val = acc.to(tl.bfloat16) # cublas-equivalent bf16 output rounding
|
||||
# implicit bf16 -> fp8 conversion on store (same as the concat-cast)
|
||||
tl.store(qout_head + offs_n[None, :], val, mask=m_mask[:, None])
|
||||
elif K_MODE == 3:
|
||||
# three-dot preload: K = 3 * BLOCK_K, a-tiles hoisted out of the N-loop.
|
||||
offs_k0 = tl.arange(0, BLOCK_K)
|
||||
offs_k1 = BLOCK_K + offs_k0
|
||||
offs_k2 = 2 * BLOCK_K + offs_k0
|
||||
a0 = tl.load(a_row + offs_k0[None, :], mask=m_mask[:, None], other=0.0)
|
||||
a1 = tl.load(a_row + offs_k1[None, :], mask=m_mask[:, None], other=0.0)
|
||||
a2 = tl.load(a_row + offs_k2[None, :], mask=m_mask[:, None], other=0.0)
|
||||
for nb in tl.static_range(N // BLOCK_N):
|
||||
offs_n = nb * BLOCK_N + tl.arange(0, BLOCK_N)
|
||||
b0 = tl.load(b_head + offs_k0[:, None] * b_s1 + offs_n[None, :] * b_s2)
|
||||
b1 = tl.load(b_head + offs_k1[:, None] * b_s1 + offs_n[None, :] * b_s2)
|
||||
b2 = tl.load(b_head + offs_k2[:, None] * b_s1 + offs_n[None, :] * b_s2)
|
||||
acc = tl.dot(a0, b0) # fp32 accumulator
|
||||
acc = tl.dot(a1, b1, acc)
|
||||
acc = tl.dot(a2, b2, acc) # same fp32 add order as the K_MODE=1 loop
|
||||
val = acc.to(tl.bfloat16) # cublas-equivalent bf16 output rounding
|
||||
# implicit bf16 -> fp8 conversion on store (same as the concat-cast)
|
||||
tl.store(qout_head + offs_n[None, :], val, mask=m_mask[:, None])
|
||||
elif K_MODE == 4:
|
||||
# padded single dot: BLOCK_K = next_pow2(K), zero-fill the k tail.
|
||||
offs_k = tl.arange(0, BLOCK_K)
|
||||
k_mask = offs_k < K
|
||||
a = tl.load(
|
||||
a_row + offs_k[None, :],
|
||||
mask=m_mask[:, None] & k_mask[None, :],
|
||||
other=0.0,
|
||||
)
|
||||
for nb in tl.static_range(N // BLOCK_N):
|
||||
offs_n = nb * BLOCK_N + tl.arange(0, BLOCK_N)
|
||||
b = tl.load(
|
||||
b_head + offs_k[:, None] * b_s1 + offs_n[None, :] * b_s2,
|
||||
mask=k_mask[:, None],
|
||||
other=0.0,
|
||||
)
|
||||
acc = tl.dot(a, b) # fp32 accumulator (padded zeros add exactly 0)
|
||||
val = acc.to(tl.bfloat16) # cublas-equivalent bf16 output rounding
|
||||
# implicit bf16 -> fp8 conversion on store (same as the concat-cast)
|
||||
tl.store(qout_head + offs_n[None, :], val, mask=m_mask[:, None])
|
||||
else:
|
||||
# K_MODE == 1: split-K loop (K % BLOCK_K == 0, e.g. K=192, BLOCK_K=64).
|
||||
for nb in tl.static_range(N // BLOCK_N):
|
||||
offs_n = nb * BLOCK_N + tl.arange(0, BLOCK_N)
|
||||
acc = tl.zeros((BLOCK_M, BLOCK_N), dtype=tl.float32)
|
||||
for kb in tl.static_range(K // BLOCK_K):
|
||||
offs_k = kb * BLOCK_K + tl.arange(0, BLOCK_K)
|
||||
a = tl.load(a_row + offs_k[None, :], mask=m_mask[:, None], other=0.0)
|
||||
b = tl.load(b_head + offs_k[:, None] * b_s1 + offs_n[None, :] * b_s2)
|
||||
acc = tl.dot(a, b, acc) # fp32 accumulator across K-blocks
|
||||
val = acc.to(tl.bfloat16) # cublas-equivalent bf16 output rounding
|
||||
# implicit bf16 -> fp8 conversion on store (same as the concat-cast)
|
||||
tl.store(qout_head + offs_n[None, :], val, mask=m_mask[:, None])
|
||||
offs_r = tl.arange(0, ROPE)
|
||||
r = tl.load(
|
||||
rope_ptr + offs_m64[:, None] * rope_s0 + h * rope_s1 + offs_r[None, :],
|
||||
mask=m_mask[:, None],
|
||||
other=0.0,
|
||||
)
|
||||
tl.store(qout_head + N + offs_r[None, :], r, mask=m_mask[:, None])
|
||||
|
||||
|
||||
# Non-power-of-2-K variant used by variant="auto" (power-of-2 K always takes
|
||||
# the single-dot fast path). Set to the winner of the K=192 A/B in
|
||||
# benchmark/kernels/deepseek/benchmark_q8kv8_q_prep.py; "loop" = the pre-A/B split-K behavior.
|
||||
_AUTO_NONPOW2_VARIANT = "two_dot"
|
||||
|
||||
|
||||
def _qprep_env_variant():
|
||||
from sglang.srt.environ import envs
|
||||
|
||||
return envs.SGLANG_OPT_Q8KV8_QPREP_VARIANT.get()
|
||||
|
||||
|
||||
# Resolved once at import (matches the module-constant style above). "auto"
|
||||
# keeps the per-K dispatch; "cuda" routes every shape to the hand-written
|
||||
# SM90 WGMMA kernel (bitwise-identical to two_dot; 1.16-1.38x faster).
|
||||
_ENV_QPREP_VARIANT = None
|
||||
|
||||
|
||||
def absorbed_bmm_concat_cast_q_fp8(
|
||||
q_fp8_pad: "torch.Tensor",
|
||||
q_nope: "torch.Tensor",
|
||||
w_kc: "torch.Tensor",
|
||||
q_rope: "torch.Tensor",
|
||||
num_heads: int,
|
||||
block_m: int = 128,
|
||||
block_n: int = 64,
|
||||
variant: str = "auto",
|
||||
block_k: int = 0,
|
||||
num_warps: int = 8,
|
||||
num_stages: int = 0,
|
||||
):
|
||||
"""Fused absorbed-q bmm + nope/rope concat + bf16->fp8 cast ("born fp8" q).
|
||||
|
||||
Replaces ``torch.bmm(q_nope.transpose(0, 1), w_kc).transpose(0, 1)``
|
||||
followed by ``concat_and_cast_q_fp8_pad`` on the Q8KV8 sparse-prefill
|
||||
path, writing the active ``[:, :num_heads, :]`` slice of the padded fp8 q
|
||||
buffer directly. Inputs:
|
||||
|
||||
* ``q_fp8_pad``: [num_tokens, pad_heads, N + ROPE] fp8_e4m3 destination.
|
||||
* ``q_nope``: [num_tokens, H, K] bf16 pre-absorb q (strided views OK).
|
||||
* ``w_kc``: [H, K, N] bf16 absorbed weight (any strides).
|
||||
* ``q_rope``: [num_tokens, H, ROPE] bf16 post-rope q (strided views OK).
|
||||
|
||||
The rope half is bit-exact vs ``concat_and_cast_q_fp8_pad``. The nope
|
||||
half keeps the same rounding stages (fp32 accum -> bf16 -> fp8) but a
|
||||
different fp32 accumulation order than cublas, so it is near- but not
|
||||
guaranteed bit-exact; keep this path behind
|
||||
``SGLANG_ENABLE_DSA_Q8KV8_BORN_FP8_Q``.
|
||||
|
||||
K (``qk_nope_head_dim``) supports any multiple of 16 in [16, 256].
|
||||
Power-of-2 K (DeepSeek 128) always takes the preload-once
|
||||
single-``tl.dot`` path. For other K (GLM 192), ``variant`` selects the
|
||||
K-dimension codegen (every variant keeps the identical fp32 -> bf16 ->
|
||||
fp8 epilogue):
|
||||
|
||||
* ``"auto"``: the current production choice (see
|
||||
``_AUTO_NONPOW2_VARIANT``).
|
||||
* ``"loop"``: split-K accumulator loop, ``BLOCK_K`` = ``block_k`` or the
|
||||
largest power-of-2 divisor of K capped at 128 (192 -> 64 x 3).
|
||||
* ``"two_dot"``: preload a as two power-of-2 tiles (192 = 128 + 64), two
|
||||
chained ``tl.dot`` per N-block, no K-loop.
|
||||
* ``"three_dot"``: preload a as three K/3 tiles (192 = 3 x 64), three
|
||||
chained ``tl.dot`` per N-block; same fp32 add order as ``"loop"``.
|
||||
* ``"pad"``: single ``tl.dot`` with ``BLOCK_K`` = next_pow2(K) (192 ->
|
||||
256) and zero-masked k tails.
|
||||
* ``"single_k"``: single ``tl.dot`` with ``BLOCK_K`` == K. Only
|
||||
compiles if the Triton build supports non-power-of-2 ``tl.arange``
|
||||
(Triton <= 3.5.x raises "arange's range must be a power of 2").
|
||||
``block_m`` / ``block_n`` / ``num_warps`` / ``num_stages`` are tuning
|
||||
knobs for the microbench sweep (0 = Triton default for ``num_stages``).
|
||||
"""
|
||||
num_tokens, _, k_dim = q_nope.shape
|
||||
n_dim = w_kc.shape[-1]
|
||||
rope_dim = q_rope.shape[-1]
|
||||
assert q_fp8_pad.dtype == torch.float8_e4m3fn
|
||||
assert q_nope.dtype == torch.bfloat16 and w_kc.dtype == torch.bfloat16
|
||||
assert q_rope.dtype == torch.bfloat16
|
||||
assert q_nope.shape[1] == num_heads and q_rope.shape[1] == num_heads
|
||||
assert w_kc.shape[0] == num_heads and w_kc.shape[1] == k_dim
|
||||
assert q_fp8_pad.shape[0] >= num_tokens and q_fp8_pad.shape[1] >= num_heads
|
||||
assert q_fp8_pad.shape[2] == n_dim + rope_dim
|
||||
# tl.arange / tl.dot constraints
|
||||
assert (
|
||||
k_dim % 16 == 0 and 16 <= k_dim <= 256
|
||||
), "K must be a multiple of 16 in [16, 256]"
|
||||
assert (rope_dim & (rope_dim - 1)) == 0, "ROPE must be a power of two"
|
||||
assert n_dim % block_n == 0, "N must be a multiple of block_n"
|
||||
assert q_nope.stride(2) == 1 and q_rope.stride(2) == 1
|
||||
assert q_fp8_pad.stride(2) == 1
|
||||
# Env override for production dispatch (SGLANG_OPT_Q8KV8_QPREP_VARIANT):
|
||||
# "auto" (default) keeps the per-K Triton dispatch; "cuda" routes every
|
||||
# shape to the WGMMA kernel below.
|
||||
global _ENV_QPREP_VARIANT
|
||||
if _ENV_QPREP_VARIANT is None:
|
||||
_ENV_QPREP_VARIANT = _qprep_env_variant()
|
||||
if variant == "auto" and _ENV_QPREP_VARIANT != "auto":
|
||||
variant = _ENV_QPREP_VARIANT
|
||||
# Hand-written SM90 WGMMA kernel (opt-in only; "auto" never routes here).
|
||||
# Same fp32 -> bf16 -> fp8 epilogue; bitwise identical to "two_dot" on
|
||||
# SM90. Requires K in {128, 192} and the production N-major w_kc layout
|
||||
# (see the wrapper's asserts); the Triton variants remain the
|
||||
# general-strides fallback.
|
||||
_valid = ("auto", "cuda", "loop", "two_dot", "three_dot", "pad", "single_k")
|
||||
if variant not in _valid:
|
||||
raise ValueError(
|
||||
f"unknown q-prep variant {variant!r} "
|
||||
f"(SGLANG_OPT_Q8KV8_QPREP_VARIANT); valid: {_valid}"
|
||||
)
|
||||
if variant == "cuda":
|
||||
from sglang.kernels.ops.attention.qprep_bf16_fp8_sm90 import q8kv8_qprep_fwd
|
||||
|
||||
q8kv8_qprep_fwd(q_fp8_pad, q_nope, w_kc, q_rope, num_heads)
|
||||
return
|
||||
# Resolve (K_MODE, BLOCK_K) from the variant; see the kernel's K-handling
|
||||
# comment for what each mode compiles to.
|
||||
if k_dim & (k_dim - 1) == 0:
|
||||
# power-of-2 K: every variant collapses to the single-dot fast path.
|
||||
k_mode, blk_k = 0, k_dim
|
||||
else:
|
||||
v = _AUTO_NONPOW2_VARIANT if variant == "auto" else variant
|
||||
if v == "loop":
|
||||
# Largest power-of-2 divisor of K, capped at 128 (K % 16 == 0
|
||||
# makes this >= 16), unless the caller pinned block_k.
|
||||
blk_k = block_k or min(k_dim & -k_dim, 128)
|
||||
assert (
|
||||
k_dim % blk_k == 0 and blk_k & (blk_k - 1) == 0 and blk_k >= 16
|
||||
), "loop needs BLOCK_K a power-of-2 divisor of K >= 16"
|
||||
k_mode = 1
|
||||
elif v == "two_dot":
|
||||
blk_k = 1 << (k_dim.bit_length() - 1) # largest power of 2 < K
|
||||
k1 = k_dim - blk_k
|
||||
assert (
|
||||
k1 & (k1 - 1) == 0 and k1 >= 16
|
||||
), "two_dot needs K = pow2 + pow2 with both halves >= 16"
|
||||
k_mode = 2
|
||||
elif v == "three_dot":
|
||||
blk_k = k_dim // 3
|
||||
assert (
|
||||
k_dim % 3 == 0 and blk_k & (blk_k - 1) == 0 and blk_k >= 16
|
||||
), "three_dot needs K = 3 * pow2 with pow2 >= 16"
|
||||
k_mode = 3
|
||||
elif v == "pad":
|
||||
blk_k = 1 << k_dim.bit_length() # next power of 2 above K
|
||||
k_mode = 4
|
||||
elif v == "single_k":
|
||||
# Non-power-of-2 BLOCK_K == K: compiles only on Triton builds
|
||||
# that allow non-power-of-2 tl.arange (not 3.5.x).
|
||||
blk_k = k_dim
|
||||
k_mode = 0
|
||||
else:
|
||||
raise ValueError(f"unknown absorbed-bmm K variant: {variant!r}")
|
||||
extra = {"num_stages": num_stages} if num_stages else {}
|
||||
grid = (triton.cdiv(num_tokens, block_m), num_heads)
|
||||
absorbed_bmm_concat_cast_q_fp8_kernel[grid](
|
||||
q_fp8_pad,
|
||||
q_nope,
|
||||
w_kc,
|
||||
q_rope,
|
||||
num_tokens,
|
||||
q_fp8_pad.stride(0),
|
||||
q_fp8_pad.stride(1),
|
||||
q_nope.stride(0),
|
||||
q_nope.stride(1),
|
||||
w_kc.stride(0),
|
||||
w_kc.stride(1),
|
||||
w_kc.stride(2),
|
||||
q_rope.stride(0),
|
||||
q_rope.stride(1),
|
||||
K=k_dim,
|
||||
N=n_dim,
|
||||
ROPE=rope_dim,
|
||||
BLOCK_M=block_m,
|
||||
BLOCK_N=block_n,
|
||||
BLOCK_K=blk_k,
|
||||
K_MODE=k_mode,
|
||||
num_warps=num_warps,
|
||||
**extra,
|
||||
)
|
||||
|
||||
|
||||
@triton.jit
|
||||
def q8kv8_topk_length_backscan_kernel(
|
||||
indices_ptr,
|
||||
out_ptr,
|
||||
stride_row,
|
||||
topk,
|
||||
BLOCK: tl.constexpr,
|
||||
):
|
||||
row = tl.program_id(0).to(tl.int64)
|
||||
base = indices_ptr + row * stride_row
|
||||
off = topk
|
||||
length = 1
|
||||
found = 0
|
||||
while (found == 0) & (off > 0):
|
||||
off -= BLOCK
|
||||
idx = off + tl.arange(0, BLOCK)
|
||||
vals = tl.load(base + idx)
|
||||
pos = tl.max(tl.where(vals >= 0, idx, -1), axis=0)
|
||||
found = tl.where(pos >= 0, 1, found)
|
||||
length = tl.where(pos >= 0, pos + 1, length)
|
||||
tl.store(out_ptr + row, length)
|
||||
|
||||
|
||||
def q8kv8_topk_length_from_indices(indices: torch.Tensor) -> torch.Tensor:
|
||||
"""Per-row valid-topk count = last non-negative position + 1 (min 1).
|
||||
|
||||
``indices``: [s_q, topk] int32 topk output whose pad slots are -1.
|
||||
Backward block scan per row: the loop exits at the first block holding a
|
||||
valid entry, so the cost is proportional to the trailing pad run — one
|
||||
block (~topk/4 elements) for rows with a full topk, which dominate long
|
||||
contexts. Semantics match the unfused ``(indices >= 0) * ramp).amax``
|
||||
derivation exactly, including all-pad rows (length 1: one pad-only block
|
||||
keeps the kernel on its clamp+mask path, contributing zero).
|
||||
"""
|
||||
s_q, topk = indices.shape
|
||||
assert indices.dtype == torch.int32 and indices.stride(1) == 1
|
||||
out = torch.empty(s_q, dtype=torch.int32, device=indices.device)
|
||||
block = 512 if topk % 512 == 0 else (256 if topk % 256 == 0 else 128)
|
||||
q8kv8_topk_length_backscan_kernel[(s_q,)](
|
||||
indices,
|
||||
out,
|
||||
indices.stride(0),
|
||||
topk,
|
||||
BLOCK=block,
|
||||
)
|
||||
return out
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Decode Context Parallel (DCP) helpers.
|
||||
#
|
||||
|
||||
@@ -4,10 +4,12 @@ from typing import Optional, Tuple
|
||||
import torch
|
||||
import triton
|
||||
|
||||
from sglang.srt.environ import envs
|
||||
from sglang.srt.utils import ceil_div, is_cuda, is_musa
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
_is_cuda = is_cuda()
|
||||
_is_musa = is_musa()
|
||||
|
||||
@@ -1539,6 +1541,25 @@ def moe_ep_deepgemm_preprocess(
|
||||
) -> Tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor]:
|
||||
# For masked grouped GEMM, shape M should be multiple of the block M (current block M: {block_m}) https://github.com/deepseek-ai/DeepGEMM/blob/main/deep_gemm/jit_kernels/m_grouped_gemm.py#L165
|
||||
m_max = (hidden_states.size(0) // 256 + 1) * 256
|
||||
if (
|
||||
envs.SGLANG_OPT_DG_MASKED_M_CAP.get()
|
||||
and not torch.cuda.is_current_stream_capturing()
|
||||
):
|
||||
# (capture guard: decode CUDA-graph capture also routes through this
|
||||
# preprocess; the D2H sync is illegal mid-capture, and decode batches
|
||||
# are small enough that the uncapped m_max is harmless there.)
|
||||
# m_max reserves capacity for ALL rank tokens in EVERY local expert:
|
||||
# the [num_local_experts, m_max, *] masked-GEMM intermediates reach
|
||||
# 7+ GiB per 32k-token chunk and OOM saturated serving. The hottest
|
||||
# expert only ever holds max(masked_m) rows, so cap the padded
|
||||
# capacity there (rounded up to the DeepGEMM block-M). Costs one
|
||||
# probe dispatch-index launch + one D2H sync per MoE layer;
|
||||
# correctness is unconditional (m_cap >= max(masked_m) by
|
||||
# construction, and the final src2dst below is built with the same
|
||||
# capped stride).
|
||||
masked_m_probe, _ = fused_moe_dispatch_index(topk_ids, num_local_experts, m_max)
|
||||
m_cap = (int(masked_m_probe.max().item()) + 255) // 256 * 256
|
||||
m_max = min(m_max, max(m_cap, 256))
|
||||
expected_m = (topk_ids.numel() - 1) // num_local_experts + 1
|
||||
|
||||
masked_m, src2dst = fused_moe_dispatch_index(topk_ids, num_local_experts, m_max)
|
||||
|
||||
@@ -768,7 +768,13 @@ def invoke_fused_moe_kernel(
|
||||
# activation block-wise fp8 quantization
|
||||
assert len(block_shape) == 2
|
||||
block_n, block_k = block_shape[0], block_shape[1]
|
||||
if _is_cuda:
|
||||
if A.dtype == torch.float8_e4m3fn:
|
||||
# Pre-quantized activation (SGLANG_OPT_MOE_QUANT_ONCE): the
|
||||
# caller already ran the per-token-group quant; A_scale holds
|
||||
# the matching scales (row- or column-major, strides are
|
||||
# passed to the kernel below).
|
||||
assert A_scale is not None
|
||||
elif _is_cuda:
|
||||
A, A_scale = sglang_per_token_group_quant_fp8(A, block_k)
|
||||
else:
|
||||
A, A_scale = per_token_group_quant_fp8(A, block_k)
|
||||
|
||||
@@ -560,6 +560,14 @@ class Envs:
|
||||
# symmetric-memory kernel), OFF elsewhere (would fall back to RCCL); override
|
||||
# explicitly to force on/off on any platform.
|
||||
SGLANG_DP_USE_REDUCE_SCATTER = EnvBool(_default_hip)
|
||||
# Quantize the variable-length DP-MoE gather payload (SGLANG_DP_USE_GATHERV
|
||||
# path, prefill/extend only) to fp8-e4m3 with per-token-group-128 scales:
|
||||
# halves the gathered hidden-state bytes over NCCL; the combine
|
||||
# (reduce_scatterv) leg stays bf16 (NCCL SUM cannot run on fp8). Lossy on
|
||||
# the wire — same group quantization the MoE expert GEMMs apply to their
|
||||
# input anyway, but router/shared-expert reads see rounded values, so this
|
||||
# stays accuracy-gated and default OFF.
|
||||
SGLANG_ENABLE_DP_GATHER_FP8 = EnvBool(False)
|
||||
SGLANG_USE_AITER_UNIFIED_ATTN = EnvBool(False)
|
||||
# Select the gate/up tile layout for AITER MoE: True -> interleave
|
||||
# (matches FlyDSL `gate_mode="interleave"` kernels), False -> separated
|
||||
@@ -689,6 +697,18 @@ class Envs:
|
||||
|
||||
# DeepGemm
|
||||
SGLANG_ENABLE_JIT_DEEPGEMM = EnvBool(True)
|
||||
# Cap the DeepGEMM masked grouped-GEMM per-expert padded capacity at
|
||||
# round_up(max(masked_m), 256) instead of round_up(rank_tokens, 256):
|
||||
# shrinks the [num_local_experts, m, *] MoE intermediates ~4x under
|
||||
# load imbalance (they otherwise OOM saturated --moe-runner-backend
|
||||
# deep_gemm serving). Costs one D2H sync per MoE layer.
|
||||
SGLANG_OPT_DG_MASKED_M_CAP = EnvBool(False)
|
||||
# Drop dp-attention MAX_LEN pad rows from MoE dispatch (StandardDispatcher
|
||||
# post-translation topk_ids -> -1): pad rows otherwise run the router on
|
||||
# stale hidden values and burn expert compute whose outputs are discarded;
|
||||
# colliding pad top-ks also inflate the DeepGEMM masked-GEMM workspace to
|
||||
# OOM at saturation. Capture-safe (reads only global_num_tokens_gpu).
|
||||
SGLANG_OPT_MASK_DP_PAD_MOE = EnvBool(False)
|
||||
SGLANG_JIT_DEEPGEMM_PRECOMPILE = EnvBool(True)
|
||||
SGLANG_JIT_DEEPGEMM_FAST_WARMUP = EnvBool(False)
|
||||
SGLANG_JIT_DEEPGEMM_COMPILE_WORKERS = EnvInt(4)
|
||||
@@ -733,6 +753,35 @@ class Envs:
|
||||
SGLANG_ENABLE_PCG_DSV2_DUAL_STREAM = EnvBool(False)
|
||||
SGLANG_DSA_TOPK_BROADCAST = EnvBool(False)
|
||||
SGLANG_DISABLE_DSA_INDEXER_FUSION = EnvBool(False)
|
||||
# Opt-in perf path for --dsa-prefill-backend flashmla_sparse_q8: fuse the
|
||||
# absorbed q bmm with the nope/rope concat + fp8 cast so q is written
|
||||
# directly in fp8 ("born fp8") and the standalone concat-cast kernel
|
||||
# disappears. Not bit-exact vs the default path (same rounding stages,
|
||||
# different GEMM accumulation order), hence default OFF until accuracy-
|
||||
# gated (oracle + full-set gsm8k).
|
||||
SGLANG_ENABLE_DSA_Q8KV8_BORN_FP8_Q = EnvBool(False)
|
||||
# Opt-in perf path for --dsa-prefill-backend flashmla_sparse_q8: pass a
|
||||
# per-row valid-topk count (derived from the trailing -1 pad run of the
|
||||
# topk indices) so the kernel skips whole pad-only topk blocks instead of
|
||||
# computing masked zero contributions. Bit-exact by construction: skipped
|
||||
# blocks contain only -1 pads, and -1 entries inside the consumed range
|
||||
# still take the in-kernel clamp+mask path.
|
||||
SGLANG_ENABLE_DSA_Q8KV8_TOPK_LENGTH = EnvBool(False)
|
||||
# Opt-in: run the born-fp8 q-prep (absorbed bmm + concat + fp8 cast,
|
||||
# ~173us/layer-call) on alt_stream underneath the DSA indexer — the two
|
||||
# chains fork independently from the q_a_layernorm output. Requires
|
||||
# SGLANG_ENABLE_DSA_Q8KV8_BORN_FP8_Q; eager-prefill-only via the born
|
||||
# predicate. Coarse per-layer join keeps the single-slot born-q buffer
|
||||
# WAR-safe.
|
||||
SGLANG_ENABLE_DSA_Q8KV8_QPREP_OVERLAP = EnvBool(False)
|
||||
# Opt-in: fuse the Q8KV8 non-prefix KV prep — cast-concat k/k_rope
|
||||
# directly into the persistent fp8 kv buffer and zero the pad band in one
|
||||
# Triton kernel (replaces bf16 _cat + copy_ cast + zero_ tail).
|
||||
SGLANG_ENABLE_DSA_Q8KV8_KV_CAT_FUSION = EnvBool(False)
|
||||
# Q8KV8 born-fp8 q-prep codegen: "auto" = per-K Triton dispatch (default);
|
||||
# "cuda" = the hand-written SM90 WGMMA kernel (bitwise identical to the
|
||||
# Triton two_dot variant, 1.16-1.38x faster across GLM/DS shapes).
|
||||
SGLANG_OPT_Q8KV8_QPREP_VARIANT = EnvStr("auto")
|
||||
|
||||
# sgl-kernel
|
||||
SGLANG_SKIP_SGL_KERNEL_VERSION_CHECK = EnvBool(False)
|
||||
@@ -1114,6 +1163,14 @@ class Envs:
|
||||
SGLANG_OPT_USE_JIT_EP_ACTIVATION = EnvBool(True)
|
||||
SGLANG_OPT_FUSE_WQA_WKV = EnvBool(True)
|
||||
SGLANG_OPT_SWIGLU_CLAMP_FUSION = EnvBool(True)
|
||||
# DeepSeek/GLM MoE (deepseek_v2.py): quantize the (dp-gathered) MoE input
|
||||
# to per-token-group-128 fp8 ONCE and feed both the fused shared-expert
|
||||
# GEMM (cutlass w8a8 linear) and the routed experts' triton fused runner,
|
||||
# instead of quantizing the same [T, hidden] tensor twice with different
|
||||
# scale layouts. Only engages on CUDA with fp8 block-128 weights, the
|
||||
# standard dispatcher, and the triton MoE runner; falls back silently
|
||||
# otherwise.
|
||||
SGLANG_OPT_MOE_QUANT_ONCE = EnvBool(False)
|
||||
|
||||
# Cache / overlap
|
||||
SGLANG_OPT_USE_FUSED_STORE_CACHE = EnvBool(True)
|
||||
|
||||
@@ -19,6 +19,7 @@ from sglang.srt.runtime_context import get_parallel
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
from sglang.kernels.ops.attention.dsa.dequant_k_cache import (
|
||||
concat_cast_kv_fp8_pad,
|
||||
dequantize_k_cache_paged,
|
||||
gather_dequant_requant_fp8_paged,
|
||||
)
|
||||
@@ -30,6 +31,7 @@ from sglang.kernels.ops.attention.dsa.transform_index import (
|
||||
from sglang.kernels.ops.attention.utils import (
|
||||
concat_mla_absorb_q_general,
|
||||
mla_quantize_and_rope_for_fp8,
|
||||
q8kv8_topk_length_from_indices,
|
||||
seqlens_expand_triton,
|
||||
)
|
||||
from sglang.kernels.ops.kvcache.cache_ops import concat_and_cast_q_fp8_pad
|
||||
@@ -495,6 +497,41 @@ class DeepseekSparseAttnBackend(
|
||||
# Q8KV8 dispatch (no-ops for other backends).
|
||||
self._q8kv8_identity_scale: Optional[torch.Tensor] = None
|
||||
self._q8kv8_qpad_buf: Optional[torch.Tensor] = None
|
||||
# Persistent (grow-only) fp8 KV destination for the Q8KV8 prefill
|
||||
# gather: [capacity_rows, 576]. Avoids a fresh torch.zeros
|
||||
# (alloc + full-buffer FillFunctor) per layer per call; only the
|
||||
# `topk` -1-sentinel landing-pad rows need zeroing each call, and
|
||||
# the gather kernel fuses that in. Same single-stream reuse
|
||||
# argument as `_q8kv8_qpad_buf`.
|
||||
self._q8kv8_kv_buf: Optional[torch.Tensor] = None
|
||||
# Per-row valid-topk early-exit (SGLANG_ENABLE_DSA_Q8KV8_TOPK_LENGTH):
|
||||
# rows whose topk indices end in a -1 pad run skip whole topk blocks
|
||||
# in-kernel.
|
||||
self._q8kv8_topk_length_enabled: bool = (
|
||||
envs.SGLANG_ENABLE_DSA_Q8KV8_TOPK_LENGTH.get()
|
||||
)
|
||||
# Persistent (grow-only) kernel-output buffers (out/max_logits/lse).
|
||||
self._q8kv8_out_bufs: Optional[tuple] = None
|
||||
# Fused non-prefix KV prep (cast-concat k/k_rope directly into the
|
||||
# fp8 buffer; SGLANG_ENABLE_DSA_Q8KV8_KV_CAT_FUSION).
|
||||
self._q8kv8_kv_cat_fusion: bool = (
|
||||
envs.SGLANG_ENABLE_DSA_Q8KV8_KV_CAT_FUSION.get()
|
||||
)
|
||||
|
||||
# Born-fp8 q handshake (SGLANG_ENABLE_DSA_Q8KV8_BORN_FP8_Q): when the
|
||||
# model's q-prep decides (via q8kv8_born_fp8_q_eligible) that this
|
||||
# batch's forward_extend is guaranteed to hit
|
||||
# _forward_flashmla_sparse_q8kv8, it writes the padded fp8 q directly
|
||||
# (fused absorbed-bmm + concat + cast) into _q8kv8_born_q_buf and
|
||||
# stashes (num_tokens, layer_id); the helper consumes the stash
|
||||
# instead of rebuilding q_fp8. Same single-stream reuse argument as
|
||||
# _q8kv8_qpad_buf. The bf16 q that flows through the attention API in
|
||||
# that mode is a NaN-poisoned sentinel: any code path that reads it by
|
||||
# mistake fails loudly instead of producing silently wrong output.
|
||||
self._q8kv8_born_q_buf: Optional[torch.Tensor] = None
|
||||
self._q8kv8_born_q_stash: Optional[Tuple[int, int]] = None
|
||||
self._q8kv8_born_q_sentinel: Optional[torch.Tensor] = None
|
||||
self._q8kv8_born_q_tbo = model_runner.server_args.enable_two_batch_overlap
|
||||
|
||||
from sglang.kernels.ops.attention.flash_mla_sm120 import (
|
||||
_validate_flashinfer_sparse_mla_backend,
|
||||
@@ -2083,6 +2120,24 @@ class DeepseekSparseAttnBackend(
|
||||
page_table_1=page_table_1,
|
||||
sm_scale=layer.scaling,
|
||||
v_head_dim=layer.v_head_dim,
|
||||
layer_id=layer.layer_id,
|
||||
)
|
||||
if self._q8kv8_kv_cat_fusion:
|
||||
# Fused path: no bf16 concat materialization — k and
|
||||
# k_rope are cast-concatenated straight into the fp8
|
||||
# buffer inside the helper.
|
||||
return self._forward_flashmla_sparse_q8kv8(
|
||||
q_nope=q_nope,
|
||||
q_rope=q_rope,
|
||||
kv_bf16=None,
|
||||
kv_k=k,
|
||||
kv_k_rope=k_rope,
|
||||
paged_kv_cache=None,
|
||||
page_table_1_flattened=None,
|
||||
page_table_1=page_table_1,
|
||||
sm_scale=layer.scaling,
|
||||
v_head_dim=layer.v_head_dim,
|
||||
layer_id=layer.layer_id,
|
||||
)
|
||||
kv_cache = _cat([k, k_rope], dim=-1)
|
||||
return self._forward_flashmla_sparse_q8kv8(
|
||||
@@ -2094,6 +2149,7 @@ class DeepseekSparseAttnBackend(
|
||||
page_table_1=page_table_1,
|
||||
sm_scale=layer.scaling,
|
||||
v_head_dim=layer.v_head_dim,
|
||||
layer_id=layer.layer_id,
|
||||
)
|
||||
|
||||
# bf16 path (dsa_impl == "flashmla_sparse").
|
||||
@@ -2428,6 +2484,97 @@ class DeepseekSparseAttnBackend(
|
||||
|
||||
return o
|
||||
|
||||
def q8kv8_born_fp8_q_eligible(
|
||||
self, forward_batch: ForwardBatch, num_heads: int
|
||||
) -> bool:
|
||||
"""True iff this batch's forward_extend is guaranteed to consume q via
|
||||
``_forward_flashmla_sparse_q8kv8`` (born-fp8 q handshake precondition).
|
||||
|
||||
Must stay in lockstep with the forward_extend dispatch: a True here
|
||||
while dispatch takes any other branch would leak the NaN sentinel into
|
||||
a real attention kernel (loud NaNs, not silent corruption, but still a
|
||||
failed forward).
|
||||
"""
|
||||
if self.dsa_prefill_impl != "flashmla_sparse_q8":
|
||||
return False
|
||||
# RAGGED routing requires exactly EXTEND (excludes decode/idle, MIXED,
|
||||
# target-verify and draft-extend, which use dsa_decode_impl anyway).
|
||||
if forward_batch.forward_mode != ForwardMode.EXTEND:
|
||||
return False
|
||||
# Per-batch dense fallback (il <= threshold) reads bf16 q directly.
|
||||
if self.use_mha:
|
||||
return False
|
||||
if self.hisparse_coordinator is not None:
|
||||
return False
|
||||
# TBO interleaves two micro-batches through one backend instance; the
|
||||
# single-slot stash handshake is not safe there.
|
||||
if self._q8kv8_born_q_tbo:
|
||||
return False
|
||||
if is_dsa_enable_prefill_cp():
|
||||
return False
|
||||
if (
|
||||
self.get_topk_transform_method(forward_batch.forward_mode)
|
||||
!= TopkTransformMethod.RAGGED
|
||||
):
|
||||
return False
|
||||
# Mirror the helper's head-padding compatibility check.
|
||||
if num_heads % 64 != 0 and 64 % num_heads != 0:
|
||||
return False
|
||||
return True
|
||||
|
||||
def q8kv8_acquire_born_q_buffer(
|
||||
self, num_tokens: int, num_heads: int, head_dim: int, device: torch.device
|
||||
) -> torch.Tensor:
|
||||
"""Padded fp8 q destination for the born-fp8 kernel (grow-only).
|
||||
|
||||
Pad rows [num_heads:pad_heads] are zeroed at allocation and never
|
||||
written afterwards (the fused kernel only writes the active heads),
|
||||
matching the _q8kv8_qpad_buf invariant the SM90 kernel relies on.
|
||||
"""
|
||||
pad = 64
|
||||
padded_heads = num_heads if num_heads % pad == 0 else pad
|
||||
buf = self._q8kv8_born_q_buf
|
||||
if (
|
||||
buf is None
|
||||
or buf.shape[0] < num_tokens
|
||||
or buf.shape[1] != padded_heads
|
||||
or buf.shape[2] != head_dim
|
||||
):
|
||||
buf = torch.zeros(
|
||||
(num_tokens, padded_heads, head_dim),
|
||||
dtype=torch.float8_e4m3fn,
|
||||
device=device,
|
||||
)
|
||||
self._q8kv8_born_q_buf = buf
|
||||
return buf[:num_tokens]
|
||||
|
||||
def q8kv8_stash_born_q(self, num_tokens: int, layer_id: int) -> None:
|
||||
if self._q8kv8_born_q_stash is not None:
|
||||
raise RuntimeError(
|
||||
"q8kv8 born-fp8 q stash was never consumed (previous stash "
|
||||
f"{self._q8kv8_born_q_stash}, new ({num_tokens}, {layer_id})): "
|
||||
"the eligibility predicate fired but forward_extend dispatched "
|
||||
"away from _forward_flashmla_sparse_q8kv8."
|
||||
)
|
||||
self._q8kv8_born_q_stash = (num_tokens, layer_id)
|
||||
|
||||
def q8kv8_born_q_sentinel(
|
||||
self, num_tokens: int, num_heads: int, v_head_dim: int, device: torch.device
|
||||
) -> torch.Tensor:
|
||||
"""NaN-poisoned bf16 stand-in for q_nope_out in born-fp8 mode.
|
||||
|
||||
Only its shape/dtype/device are ever legitimately used downstream; a
|
||||
NaN payload turns any accidental read into loud NaN output.
|
||||
"""
|
||||
numel = num_tokens * num_heads * v_head_dim
|
||||
buf = self._q8kv8_born_q_sentinel
|
||||
if buf is None or buf.numel() < numel:
|
||||
buf = torch.full(
|
||||
(numel,), float("nan"), dtype=torch.bfloat16, device=device
|
||||
)
|
||||
self._q8kv8_born_q_sentinel = buf
|
||||
return buf[:numel].view(num_tokens, num_heads, v_head_dim)
|
||||
|
||||
def _forward_flashmla_sparse_q8kv8(
|
||||
self,
|
||||
q_nope: torch.Tensor,
|
||||
@@ -2438,6 +2585,9 @@ class DeepseekSparseAttnBackend(
|
||||
sm_scale: float,
|
||||
paged_kv_cache: Optional[torch.Tensor] = None,
|
||||
page_table_1_flattened: Optional[torch.Tensor] = None,
|
||||
layer_id: Optional[int] = None,
|
||||
kv_k: Optional[torch.Tensor] = None,
|
||||
kv_k_rope: Optional[torch.Tensor] = None,
|
||||
) -> torch.Tensor:
|
||||
"""Native FP8 (q8 x kv8) sparse-prefill attention (SM90 JIT kernel).
|
||||
|
||||
@@ -2472,12 +2622,37 @@ class DeepseekSparseAttnBackend(
|
||||
required_padding = 64
|
||||
need_padding = num_heads % required_padding != 0
|
||||
|
||||
# Born-fp8 fast path (SGLANG_ENABLE_DSA_Q8KV8_BORN_FP8_Q): the model's
|
||||
# q-prep already wrote the padded fp8 q (fused absorbed-bmm + concat +
|
||||
# cast); consume the stash instead of rebuilding it. q_nope here is
|
||||
# the NaN sentinel (shape-only); q_rope's bf16 content is valid but
|
||||
# unused.
|
||||
born = self._q8kv8_born_q_stash
|
||||
if born is not None:
|
||||
self._q8kv8_born_q_stash = None
|
||||
born_tokens, born_layer_id = born
|
||||
if born_tokens != num_tokens or (
|
||||
layer_id is not None and born_layer_id != layer_id
|
||||
):
|
||||
raise RuntimeError(
|
||||
"q8kv8 born-fp8 q stash mismatch: stashed "
|
||||
f"(num_tokens={born_tokens}, layer_id={born_layer_id}) but "
|
||||
f"consuming (num_tokens={num_tokens}, layer_id={layer_id})."
|
||||
)
|
||||
q_fp8 = self._q8kv8_born_q_buf[:num_tokens]
|
||||
expected_heads = required_padding if need_padding else num_heads
|
||||
if q_fp8.shape[1] != expected_heads or q_fp8.shape[2] != head_dim:
|
||||
raise RuntimeError(
|
||||
"q8kv8 born-fp8 q buffer shape mismatch: got "
|
||||
f"{tuple(q_fp8.shape)}, expected (*, {expected_heads}, "
|
||||
f"{head_dim})."
|
||||
)
|
||||
# Build the fp8 q. concat_and_cast_q_fp8_pad fuses the nope/rope
|
||||
# concat with the bf16->fp8 cast in one Triton kernel (bit-exact vs
|
||||
# concat + .to(fp8)); it requires power-of-two head/dim counts (a
|
||||
# tl.arange constraint), so non-power-of-two head counts fall back to
|
||||
# the generic concat + cast.
|
||||
if need_padding:
|
||||
elif need_padding:
|
||||
if required_padding % num_heads != 0:
|
||||
raise ValueError(
|
||||
f"num_heads={num_heads} cannot be padded to {required_padding}; "
|
||||
@@ -2521,21 +2696,85 @@ class DeepseekSparseAttnBackend(
|
||||
# Mapping many slots onto one shared row would serialize the kernel's
|
||||
# KV gather; distinct zero rows are value-identical (zero KV
|
||||
# contributes nothing to the softmax-weighted sum) at full speed.
|
||||
#
|
||||
# The destination is a persistent grow-only buffer instead of a fresh
|
||||
# torch.zeros: rows [0, num_kv_tokens) are fully overwritten every
|
||||
# call (gather kernel / cast-copy), so only the pad rows
|
||||
# [num_kv_tokens, num_kv_tokens + topk) - exactly the rows the SM90
|
||||
# kernel's -1 clamp (pad_base + slot) can read - need zeroing, and
|
||||
# they need it EVERY call because a previous, larger call may have
|
||||
# left real KV data there. The gather kernel fuses the pad-row
|
||||
# zeroing; the bf16 path zeroes the tail explicitly.
|
||||
topk = page_table_1.shape[-1]
|
||||
if paged_kv_cache is not None:
|
||||
num_kv_tokens = page_table_1_flattened.shape[0]
|
||||
elif kv_k is not None:
|
||||
num_kv_tokens = kv_k.shape[0]
|
||||
else:
|
||||
num_kv_tokens = kv_bf16.shape[0]
|
||||
total_kv_rows = num_kv_tokens + topk
|
||||
kv_buf = self._q8kv8_kv_buf
|
||||
if kv_buf is None or kv_buf.shape[0] < total_kv_rows:
|
||||
kv_buf = torch.empty(
|
||||
(total_kv_rows, head_dim),
|
||||
dtype=torch.float8_e4m3fn,
|
||||
device=dev,
|
||||
)
|
||||
self._q8kv8_kv_buf = kv_buf
|
||||
if paged_kv_cache is not None:
|
||||
kv_padded = gather_dequant_requant_fp8_paged(
|
||||
paged_kv_cache,
|
||||
page_table_1_flattened,
|
||||
extra_rows=topk,
|
||||
out=kv_buf[:total_kv_rows],
|
||||
).view(-1, 1, head_dim)
|
||||
elif kv_k is not None:
|
||||
# Fused non-prefix KV prep (SGLANG_ENABLE_DSA_Q8KV8_KV_CAT_FUSION):
|
||||
# cast-concat k/k_rope straight into the fp8 buffer + zero the pad
|
||||
# band in ONE kernel — the bf16 _cat materialization, the copy_
|
||||
# cast and the zero_ tail all disappear. Same store-cast as the
|
||||
# gather kernel (bit-identical bytes).
|
||||
kv_padded = concat_cast_kv_fp8_pad(
|
||||
kv_buf[:total_kv_rows], kv_k, kv_k_rope, num_kv_tokens
|
||||
).view(-1, 1, head_dim)
|
||||
else:
|
||||
kv_padded = kv_bf16.new_zeros(
|
||||
(kv_bf16.shape[0] + topk, *kv_bf16.shape[1:]),
|
||||
dtype=torch.float8_e4m3fn,
|
||||
)
|
||||
kv_padded[: kv_bf16.shape[0]].copy_(kv_bf16)
|
||||
kv_padded = kv_buf[:total_kv_rows]
|
||||
# bf16 -> fp8 cast copy, same op as the previous fresh-buffer
|
||||
# path (bit-identical bytes).
|
||||
kv_padded[:num_kv_tokens].copy_(kv_bf16.view(num_kv_tokens, head_dim))
|
||||
kv_padded[num_kv_tokens:].zero_()
|
||||
kv_padded = kv_padded.view(-1, 1, head_dim)
|
||||
|
||||
# Per-row valid-topk count = last non-pad position + 1. Bit-exact
|
||||
# vs topk_length=None: the skipped tail blocks contain only -1 pads
|
||||
# (masked to zero contribution today), and -1 entries inside the
|
||||
# consumed range still take the kernel's clamp+mask path. The
|
||||
# backscan's cost is proportional to the trailing pad run, so rows
|
||||
# with a full topk (all rows at long context) pay ~one block read.
|
||||
topk_length = None
|
||||
if self._q8kv8_topk_length_enabled:
|
||||
topk_length = q8kv8_topk_length_from_indices(page_table_1)
|
||||
|
||||
# Persistent kernel-output buffers (out / max_logits / lse): the
|
||||
# wrapper otherwise torch.empty's all three per layer-call. The
|
||||
# kernel fully overwrites the active [:s_q] rows and everything runs
|
||||
# on one stream, so reuse is safe — same argument as _q8kv8_qpad_buf.
|
||||
s_q, pad_heads = q_fp8.shape[0], q_fp8.shape[1]
|
||||
out_bufs = self._q8kv8_out_bufs
|
||||
if (
|
||||
out_bufs is None
|
||||
or out_bufs[0].shape[0] < s_q
|
||||
or out_bufs[0].shape[1] != pad_heads
|
||||
):
|
||||
out_bufs = (
|
||||
torch.empty(
|
||||
s_q, pad_heads, v_head_dim, dtype=torch.bfloat16, device=dev
|
||||
),
|
||||
torch.empty(s_q, pad_heads, dtype=torch.float32, device=dev),
|
||||
torch.empty(s_q, pad_heads, dtype=torch.float32, device=dev),
|
||||
)
|
||||
self._q8kv8_out_bufs = out_bufs
|
||||
|
||||
o, _, _ = sparse_mla_q8kv8_prefill_fwd(
|
||||
q=q_fp8,
|
||||
kv=kv_padded,
|
||||
@@ -2545,7 +2784,10 @@ class DeepseekSparseAttnBackend(
|
||||
kv_scale=identity_scale,
|
||||
d_v=v_head_dim,
|
||||
attn_sink=None,
|
||||
topk_length=None,
|
||||
topk_length=topk_length,
|
||||
out=out_bufs[0][:s_q],
|
||||
max_logits=out_bufs[1][:s_q],
|
||||
lse=out_bufs[2][:s_q],
|
||||
)
|
||||
|
||||
# Trim the output back to the original head count if we padded.
|
||||
|
||||
@@ -1,11 +1,14 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import functools
|
||||
import logging
|
||||
from contextlib import contextmanager
|
||||
from enum import IntEnum, auto
|
||||
from typing import TYPE_CHECKING, List, Optional, Tuple
|
||||
|
||||
import torch
|
||||
import triton
|
||||
import triton.language as tl
|
||||
|
||||
from sglang.srt.distributed import (
|
||||
GroupCoordinator,
|
||||
@@ -136,6 +139,7 @@ class _DpGatheredBufferWrapper:
|
||||
_local_dp_buffer_len: int = 0
|
||||
_dp_max_padding: bool = False
|
||||
_global_num_tokens: Optional[List[int]] = None
|
||||
_global_num_tokens_gpu: Optional[torch.Tensor] = None
|
||||
|
||||
@classmethod
|
||||
def set_metadata(cls, hidden_size: int, dtype: torch.dtype, device: torch.device):
|
||||
@@ -153,11 +157,13 @@ class _DpGatheredBufferWrapper:
|
||||
local_dp_buffer_len: int,
|
||||
dp_max_padding: bool,
|
||||
global_num_tokens: Optional[List[int]] = None,
|
||||
global_num_tokens_gpu: Optional[torch.Tensor] = None,
|
||||
):
|
||||
cls._global_dp_buffer_len = global_dp_buffer_len
|
||||
cls._local_dp_buffer_len = local_dp_buffer_len
|
||||
cls._dp_max_padding = dp_max_padding
|
||||
cls._global_num_tokens = global_num_tokens
|
||||
cls._global_num_tokens_gpu = global_num_tokens_gpu
|
||||
|
||||
@classmethod
|
||||
def get_global_dp_buffer(cls, group: GroupCoordinator) -> torch.Tensor:
|
||||
@@ -201,6 +207,10 @@ class _DpGatheredBufferWrapper:
|
||||
def get_dp_global_num_tokens(cls) -> List[int]:
|
||||
return cls._global_num_tokens
|
||||
|
||||
@classmethod
|
||||
def get_dp_global_num_tokens_gpu(cls) -> Optional[torch.Tensor]:
|
||||
return cls._global_num_tokens_gpu
|
||||
|
||||
@classmethod
|
||||
def get_dp_hidden_size(cls) -> int:
|
||||
from sglang.srt.runtime_context import get_flags
|
||||
@@ -229,9 +239,14 @@ def set_dp_buffer_len(
|
||||
local_dp_buffer_len: int,
|
||||
dp_max_padding: bool,
|
||||
global_num_tokens: Optional[List[int]] = None,
|
||||
global_num_tokens_gpu: Optional[torch.Tensor] = None,
|
||||
):
|
||||
_DpGatheredBufferWrapper.set_dp_buffer_len(
|
||||
global_dp_buffer_len, local_dp_buffer_len, dp_max_padding, global_num_tokens
|
||||
global_dp_buffer_len,
|
||||
local_dp_buffer_len,
|
||||
dp_max_padding,
|
||||
global_num_tokens,
|
||||
global_num_tokens_gpu,
|
||||
)
|
||||
|
||||
|
||||
@@ -505,6 +520,142 @@ def _dp_gather_via_all_gather(
|
||||
# tp_size==dp_size (attn_tp_size==1) case is supported for now (e.g. tp8dp8).
|
||||
_USE_DP_GATHERV = get_bool_env_var("SGLANG_DP_USE_GATHERV")
|
||||
|
||||
_DP_GATHER_FP8_GROUP = 128
|
||||
# Grow-only gathered fp8 payload / scales buffers, keyed by device.
|
||||
_dp_gather_fp8_bufs: dict = {}
|
||||
|
||||
|
||||
@functools.lru_cache(maxsize=1)
|
||||
def _use_dp_gather_fp8() -> bool:
|
||||
from sglang.srt.environ import envs
|
||||
|
||||
return envs.SGLANG_ENABLE_DP_GATHER_FP8.get()
|
||||
|
||||
|
||||
def _get_dp_gather_fp8_bufs(rows: int, hidden: int, device: torch.device):
|
||||
key = str(device)
|
||||
bufs = _dp_gather_fp8_bufs.get(key)
|
||||
if bufs is None or bufs[0].shape[0] < rows:
|
||||
bufs = (
|
||||
torch.empty((rows, hidden), dtype=torch.uint8, device=device),
|
||||
torch.empty(
|
||||
(rows, hidden // _DP_GATHER_FP8_GROUP),
|
||||
dtype=torch.float32,
|
||||
device=device,
|
||||
),
|
||||
)
|
||||
_dp_gather_fp8_bufs[key] = bufs
|
||||
return bufs[0][:rows], bufs[1][:rows]
|
||||
|
||||
|
||||
@triton.jit
|
||||
def _dequant_per_token_group_fp8_kernel(
|
||||
q_ptr,
|
||||
s_ptr,
|
||||
out_ptr,
|
||||
HIDDEN: tl.constexpr,
|
||||
NGROUPS: tl.constexpr,
|
||||
GROUP: tl.constexpr,
|
||||
BLOCK: tl.constexpr,
|
||||
):
|
||||
row = tl.program_id(0).to(tl.int64)
|
||||
# HIDDEN may not be a multiple of BLOCK (e.g. DeepSeek 7168 vs BLOCK
|
||||
# 2048): the tail iteration must be masked or it reads/writes up to
|
||||
# BLOCK-1 elements past the row (cross-row corruption + OOB on the last
|
||||
# row). HIDDEN is constexpr, so the mask folds away when it divides.
|
||||
for start in tl.static_range(0, HIDDEN, BLOCK):
|
||||
offs = start + tl.arange(0, BLOCK)
|
||||
mask = offs < HIDDEN
|
||||
qv = tl.load(q_ptr + row * HIDDEN + offs, mask=mask, other=0.0).to(tl.float32)
|
||||
sv = tl.load(s_ptr + row * NGROUPS + offs // GROUP, mask=mask, other=0.0)
|
||||
tl.store(out_ptr + row * HIDDEN + offs, (qv * sv).to(tl.bfloat16), mask=mask)
|
||||
|
||||
|
||||
@triton.jit
|
||||
def _mask_dp_pad_topk_ids_kernel(
|
||||
topk_ids_ptr,
|
||||
counts_ptr,
|
||||
max_len,
|
||||
TOPK: tl.constexpr,
|
||||
BLOCK: tl.constexpr,
|
||||
):
|
||||
row = tl.program_id(0).to(tl.int64)
|
||||
rank = row // max_len
|
||||
pos = row % max_len
|
||||
valid = pos < tl.load(counts_ptr + rank)
|
||||
if valid == 0:
|
||||
offs = tl.arange(0, BLOCK)
|
||||
tl.store(topk_ids_ptr + row * TOPK + offs, -1, mask=offs < TOPK)
|
||||
|
||||
|
||||
def mask_dp_pad_moe_topk_ids(topk_ids: torch.Tensor) -> None:
|
||||
"""Set MAX_LEN pad rows' (post-translation, local) topk_ids to -1 in place.
|
||||
|
||||
Under dp-attention MAX_LEN padding the gathered MoE buffer is
|
||||
[dp_size * max_len, hidden] with rank r's real rows at
|
||||
[r*max_len, r*max_len + global_num_tokens[r]); the pad rows carry stale
|
||||
hidden values, run the router, and get dispatched into experts whose
|
||||
outputs are then discarded by the post-reorder scatter — pure wasted
|
||||
compute, and a masked-grouped-GEMM workspace blow-up when they collide
|
||||
on the same top-k. -1 is the drop sentinel both the triton fused_moe
|
||||
(filter_expert) and the DeepGEMM EP preprocess honor; it must be applied
|
||||
AFTER the local_expert_mapping gather (a pre-translation -1 aliases to
|
||||
the mapping table's last entry). Capture-safe: per-batch state is read
|
||||
only from the replay-updated global_num_tokens_gpu tensor.
|
||||
"""
|
||||
counts = _DpGatheredBufferWrapper.get_dp_global_num_tokens_gpu()
|
||||
if counts is None:
|
||||
return
|
||||
max_len = _DpGatheredBufferWrapper.get_local_dp_buffer_len()
|
||||
rows, topk = topk_ids.shape
|
||||
if max_len <= 0 or rows != counts.shape[0] * max_len:
|
||||
# Layout mismatch (e.g. non-DP or logits-path caller): do nothing.
|
||||
return
|
||||
_mask_dp_pad_topk_ids_kernel[(rows,)](
|
||||
topk_ids,
|
||||
counts,
|
||||
max_len,
|
||||
TOPK=topk,
|
||||
BLOCK=triton.next_power_of_2(topk),
|
||||
)
|
||||
|
||||
|
||||
def _dp_gather_via_all_gatherv_fp8(
|
||||
global_tokens: torch.Tensor,
|
||||
local_real: torch.Tensor,
|
||||
sizes: List[int],
|
||||
):
|
||||
"""fp8 wire format for the variable-length DP gather: quantize the local
|
||||
rows per-token-group (the SAME group-128 quantization the MoE expert GEMMs
|
||||
apply to their input downstream), gather payload (as uint8 — NCCL has no
|
||||
fp8 dtype; the gatherv leg is broadcast-only so a byte view is safe) and
|
||||
scales in two output-buffered gatherv calls, then dequantize into the
|
||||
bf16 global buffer. Zero pad rows quantize to (q=0, s=eps) and so
|
||||
dequantize back to exact zeros — the MoE-tail invariant is preserved.
|
||||
The combine leg (reduce_scatterv) stays bf16: NCCL SUM cannot run on fp8."""
|
||||
from sglang.kernels.ops.quantization.fp8_kernel import (
|
||||
sglang_per_token_group_quant_fp8,
|
||||
)
|
||||
|
||||
rows = global_tokens.shape[0]
|
||||
hidden = global_tokens.shape[-1]
|
||||
q, s = sglang_per_token_group_quant_fp8(
|
||||
local_real.contiguous(), _DP_GATHER_FP8_GROUP
|
||||
)
|
||||
gq, gs = _get_dp_gather_fp8_bufs(rows, hidden, global_tokens.device)
|
||||
tp_group = get_tp_group()
|
||||
tp_group.all_gatherv(q.view(torch.uint8), sizes=sizes, output=gq)
|
||||
tp_group.all_gatherv(s, sizes=sizes, output=gs)
|
||||
_dequant_per_token_group_fp8_kernel[(rows,)](
|
||||
gq.view(torch.float8_e4m3fn),
|
||||
gs,
|
||||
global_tokens,
|
||||
HIDDEN=hidden,
|
||||
NGROUPS=hidden // _DP_GATHER_FP8_GROUP,
|
||||
GROUP=_DP_GATHER_FP8_GROUP,
|
||||
BLOCK=2048,
|
||||
)
|
||||
|
||||
|
||||
def is_dp_gatherv_active() -> bool:
|
||||
"""Variable-length DP-MoE gather/scatter (all_gatherv + reduce_scatterv) is
|
||||
@@ -568,6 +719,19 @@ def _dp_gather_via_all_gatherv(
|
||||
# falls back to all_reduce). Pass global_tokens as the NCCL output buffer so
|
||||
# the gather writes directly into it -- avoids the previous extra full-buffer
|
||||
# torch.cat + copy_ (two ~sum(sizes)*hidden DtoD copies, ~700us/layer at c512).
|
||||
# NOTE: the fp8 branch condition must be identical on EVERY DP rank (all
|
||||
# ranks must issue the same NCCL op sequence) — env/dtype/hidden are
|
||||
# rank-uniform; never gate on per-rank state like forward_mode (ranks can
|
||||
# be extend/idle-mixed within one global forward). Prefill-only is
|
||||
# already structural: the gatherv path runs only under SUM_LEN padding,
|
||||
# which decode-only steps and CUDA-graph capture never select.
|
||||
if (
|
||||
_use_dp_gather_fp8()
|
||||
and global_tokens.dtype == torch.bfloat16
|
||||
and global_tokens.shape[-1] % _DP_GATHER_FP8_GROUP == 0
|
||||
):
|
||||
_dp_gather_via_all_gatherv_fp8(global_tokens, local_real, sizes)
|
||||
return
|
||||
get_tp_group().all_gatherv(local_real, sizes=sizes, output=global_tokens)
|
||||
|
||||
|
||||
|
||||
@@ -1332,7 +1332,12 @@ class FusedMoE(torch.nn.Module):
|
||||
f"Unsupported weight_name {weight_name} for FusedMoE weight_loader_fused. Nothing is loaded."
|
||||
)
|
||||
|
||||
def forward(self, hidden_states: torch.Tensor, topk_output: TopKOutput):
|
||||
def forward(
|
||||
self,
|
||||
hidden_states: torch.Tensor,
|
||||
topk_output: TopKOutput,
|
||||
pre_quant_input: Optional[Tuple[torch.Tensor, torch.Tensor]] = None,
|
||||
):
|
||||
if self._use_ascend_fuseep:
|
||||
from sglang.srt.hardware_backend.npu.moe.fuseep import forward_fuseep
|
||||
|
||||
@@ -1360,11 +1365,20 @@ class FusedMoE(torch.nn.Module):
|
||||
)
|
||||
else:
|
||||
# Make sure there is torch lib op registration for the whole moe layer
|
||||
return self.forward_impl(hidden_states, topk_output)
|
||||
return self.forward_impl(
|
||||
hidden_states, topk_output, pre_quant_input=pre_quant_input
|
||||
)
|
||||
else:
|
||||
return self.forward_impl(hidden_states, topk_output)
|
||||
return self.forward_impl(
|
||||
hidden_states, topk_output, pre_quant_input=pre_quant_input
|
||||
)
|
||||
|
||||
def forward_impl(self, hidden_states: torch.Tensor, topk_output: TopKOutput):
|
||||
def forward_impl(
|
||||
self,
|
||||
hidden_states: torch.Tensor,
|
||||
topk_output: TopKOutput,
|
||||
pre_quant_input: Optional[Tuple[torch.Tensor, torch.Tensor]] = None,
|
||||
):
|
||||
origin_hidden_states_dim = hidden_states.shape[-1]
|
||||
assert self.quant_method is not None
|
||||
|
||||
@@ -1375,6 +1389,18 @@ class FusedMoE(torch.nn.Module):
|
||||
dispatch_output = self.dispatcher.dispatch(
|
||||
hidden_states=hidden_states, topk_output=topk_output
|
||||
)
|
||||
if (
|
||||
pre_quant_input is not None
|
||||
and dispatch_output.format.is_standard()
|
||||
and dispatch_output.hidden_states_scale is None
|
||||
):
|
||||
# SGLANG_OPT_MOE_QUANT_ONCE: the standard dispatch was a pure
|
||||
# passthrough, so the caller's pre-quantized (q, scale) pair still
|
||||
# matches dispatch_output.hidden_states; attach it for the triton
|
||||
# fused runner to skip its own activation quant.
|
||||
dispatch_output = dispatch_output._replace(
|
||||
hidden_states_pre_quant=pre_quant_input
|
||||
)
|
||||
|
||||
combine_input = self.run_moe_core(
|
||||
dispatch_output=dispatch_output,
|
||||
|
||||
@@ -1,5 +1,6 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import logging
|
||||
from dataclasses import dataclass
|
||||
from typing import TYPE_CHECKING, Any, List, Optional, Tuple
|
||||
|
||||
@@ -8,6 +9,9 @@ import torch
|
||||
|
||||
from sglang.kernels.ops.attention.dsv4 import silu_and_mul_masked_post_quant
|
||||
from sglang.kernels.ops.quantization import per_token_group_quant
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
from sglang.srt.distributed import get_tp_group
|
||||
from sglang.srt.distributed.device_communicators.pynccl_allocator import (
|
||||
use_symmetric_memory,
|
||||
@@ -448,9 +452,20 @@ class DeepGemmRunnerCore(MoeRunnerCore):
|
||||
|
||||
num_groups, m, k = hidden_states.shape
|
||||
n = w13_weight.size(1)
|
||||
gateup_output = torch.empty(
|
||||
(num_groups, m, n), device=hidden_states_device, dtype=torch.bfloat16
|
||||
)
|
||||
try:
|
||||
gateup_output = torch.empty(
|
||||
(num_groups, m, n), device=hidden_states_device, dtype=torch.bfloat16
|
||||
)
|
||||
except torch.OutOfMemoryError:
|
||||
logger.error(
|
||||
"Masked grouped-GEMM workspace allocation failed "
|
||||
"(num_groups=%d m=%d n=%d). If this happens under saturated "
|
||||
"dp-attention prefill, try SGLANG_OPT_DG_MASKED_M_CAP=1.",
|
||||
num_groups,
|
||||
m,
|
||||
n,
|
||||
)
|
||||
raise
|
||||
deep_gemm_wrapper.grouped_gemm_nt_f8f8bf16_masked(
|
||||
(hidden_states, hidden_states_scale),
|
||||
(w13_weight, w13_scale),
|
||||
|
||||
@@ -217,6 +217,15 @@ def fused_experts_none_to_triton(
|
||||
fused_experts,
|
||||
)
|
||||
|
||||
# SGLANG_OPT_MOE_QUANT_ONCE: use the caller's pre-quantized activation
|
||||
# (per-token-group-128 fp8 q + scales) instead of re-quantizing inside
|
||||
# invoke_fused_moe_kernel.
|
||||
pre_quant = dispatch_output.hidden_states_pre_quant
|
||||
if pre_quant is not None:
|
||||
a1_q, a1_scale = pre_quant
|
||||
else:
|
||||
a1_q, a1_scale = None, quant_info.a13_scale
|
||||
|
||||
output = fused_experts(
|
||||
hidden_states=dispatch_output.hidden_states,
|
||||
w1=quant_info.w13_weight,
|
||||
@@ -234,9 +243,10 @@ def fused_experts_none_to_triton(
|
||||
w2_scale=quant_info.w2_scale,
|
||||
w1_zp=quant_info.w13_zp,
|
||||
w2_zp=quant_info.w2_zp,
|
||||
a1_scale=quant_info.a13_scale,
|
||||
a1_scale=a1_scale,
|
||||
a2_scale=quant_info.a2_scale,
|
||||
block_shape=quant_info.block_shape,
|
||||
a1_q=a1_q,
|
||||
)
|
||||
|
||||
return StandardCombineInput(
|
||||
|
||||
@@ -128,6 +128,7 @@ def inplace_fused_experts(
|
||||
filter_expert: bool = True,
|
||||
swiglu_limit: Optional[float] = None,
|
||||
gate_up_interleaved: bool = True,
|
||||
a1_q: Optional[torch.Tensor] = None,
|
||||
) -> None:
|
||||
fused_experts_impl(
|
||||
hidden_states,
|
||||
@@ -160,6 +161,7 @@ def inplace_fused_experts(
|
||||
filter_expert,
|
||||
swiglu_limit=swiglu_limit,
|
||||
gate_up_interleaved=gate_up_interleaved,
|
||||
a1_q=a1_q,
|
||||
)
|
||||
|
||||
|
||||
@@ -194,6 +196,7 @@ def outplace_fused_experts(
|
||||
filter_expert: bool = True,
|
||||
swiglu_limit: Optional[float] = None,
|
||||
gate_up_interleaved: bool = True,
|
||||
a1_q: Optional[torch.Tensor] = None,
|
||||
) -> torch.Tensor:
|
||||
return fused_experts_impl(
|
||||
hidden_states,
|
||||
@@ -226,6 +229,7 @@ def outplace_fused_experts(
|
||||
filter_expert=filter_expert,
|
||||
swiglu_limit=swiglu_limit,
|
||||
gate_up_interleaved=gate_up_interleaved,
|
||||
a1_q=a1_q,
|
||||
)
|
||||
|
||||
|
||||
@@ -249,6 +253,7 @@ def fused_experts(
|
||||
a1_scale: Optional[torch.Tensor] = None,
|
||||
a2_scale: Optional[torch.Tensor] = None,
|
||||
block_shape: Optional[List[int]] = None,
|
||||
a1_q: Optional[torch.Tensor] = None,
|
||||
):
|
||||
topk_weights, topk_ids, _ = topk_output
|
||||
filter_expert = (
|
||||
@@ -286,6 +291,7 @@ def fused_experts(
|
||||
filter_expert,
|
||||
swiglu_limit=moe_runner_config.swiglu_limit,
|
||||
gate_up_interleaved=moe_runner_config.gate_up_interleaved,
|
||||
a1_q=a1_q,
|
||||
)
|
||||
return hidden_states
|
||||
else:
|
||||
@@ -319,6 +325,7 @@ def fused_experts(
|
||||
filter_expert=filter_expert,
|
||||
swiglu_limit=moe_runner_config.swiglu_limit,
|
||||
gate_up_interleaved=moe_runner_config.gate_up_interleaved,
|
||||
a1_q=a1_q,
|
||||
)
|
||||
|
||||
|
||||
@@ -461,12 +468,19 @@ def _fused_moe_kernel_sequence(
|
||||
hooks: Optional[Any] = None,
|
||||
swiglu_limit: Optional[float] = None,
|
||||
gate_up_interleaved: bool = True,
|
||||
a1_q: Optional[torch.Tensor] = None,
|
||||
) -> torch.Tensor:
|
||||
"""Run the MoE kernel/activation/kernel/combine sequence in a single shot.
|
||||
|
||||
Inputs are already aligned and the block-size config is already resolved.
|
||||
Supports optional LoRA hooks that fire between the two kernels and before
|
||||
combine. Returns ``out_hidden_states``.
|
||||
|
||||
``a1_q`` (SGLANG_OPT_MOE_QUANT_ONCE): optional pre-quantized fp8 view of
|
||||
``hidden_states`` for the gate-up GEMM (per-token-group ``block_shape[1]``
|
||||
quant, ``a1_scale`` holds the matching scales, rows may exceed
|
||||
``num_tokens`` due to 4-row padding). ``hidden_states`` stays bf16 and is
|
||||
still used for output dtype/shape and the inplace combine.
|
||||
"""
|
||||
num_tokens = hidden_states.shape[0]
|
||||
E, N, _ = w1.shape
|
||||
@@ -479,6 +493,17 @@ def _fused_moe_kernel_sequence(
|
||||
if hooks and (hooks.after_gate_up is not None or hooks.after_down is not None):
|
||||
down_moe_use_tma = False
|
||||
|
||||
if a1_q is not None:
|
||||
assert (
|
||||
use_fp8_w8a8
|
||||
and block_shape is not None
|
||||
and a1_scale is not None
|
||||
and a1_q.dtype == torch.float8_e4m3fn
|
||||
and a1_q.is_contiguous()
|
||||
and a1_q.shape[0] >= num_tokens
|
||||
and a1_q.shape[1] == hidden_states.shape[1]
|
||||
), "a1_q requires block-wise fp8 with matching pre-quantized activation"
|
||||
|
||||
padded_tokens = (
|
||||
min(num_tokens * topk, E + 1) * (config["BLOCK_SIZE_M"] - 1)
|
||||
if down_moe_use_tma
|
||||
@@ -520,7 +545,7 @@ def _fused_moe_kernel_sequence(
|
||||
)
|
||||
|
||||
invoke_fused_moe_kernel(
|
||||
hidden_states,
|
||||
a1_q if a1_q is not None else hidden_states,
|
||||
w1,
|
||||
b1,
|
||||
intermediate_cache1,
|
||||
@@ -866,6 +891,7 @@ def fused_experts_impl(
|
||||
filter_expert: bool = True,
|
||||
swiglu_limit: Optional[float] = None,
|
||||
gate_up_interleaved: bool = True,
|
||||
a1_q: Optional[torch.Tensor] = None,
|
||||
):
|
||||
padded_size = padding_size
|
||||
if not (use_fp8_w8a8 or use_int8_w8a8) or block_shape is not None or _use_aiter:
|
||||
@@ -942,6 +968,7 @@ def fused_experts_impl(
|
||||
hooks=None,
|
||||
swiglu_limit=swiglu_limit,
|
||||
gate_up_interleaved=gate_up_interleaved,
|
||||
a1_q=a1_q,
|
||||
)
|
||||
|
||||
|
||||
|
||||
@@ -1,6 +1,6 @@
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import TYPE_CHECKING, NamedTuple, Optional
|
||||
from typing import TYPE_CHECKING, NamedTuple, Optional, Tuple
|
||||
|
||||
import torch
|
||||
|
||||
@@ -14,6 +14,8 @@ from sglang.srt.layers.dp_attention import (
|
||||
get_dp_global_num_tokens,
|
||||
get_local_dp_buffer,
|
||||
is_allocation_symmetric,
|
||||
is_dp_max_padding,
|
||||
mask_dp_pad_moe_topk_ids,
|
||||
)
|
||||
from sglang.srt.layers.moe.moe_runner.base import MoeRunnerConfig
|
||||
from sglang.srt.layers.moe.token_dispatcher.base import (
|
||||
@@ -39,6 +41,10 @@ from sglang.srt.utils.common import (
|
||||
_is_hip = is_hip()
|
||||
_use_aiter = get_bool_env_var("SGLANG_USE_AITER") and _is_hip
|
||||
|
||||
from sglang.srt.environ import envs as _envs
|
||||
|
||||
_MASK_DP_PAD_MOE = _envs.SGLANG_OPT_MASK_DP_PAD_MOE.get()
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from sglang.srt.layers.moe.topk import TopKOutput
|
||||
|
||||
@@ -62,6 +68,11 @@ class StandardDispatchOutput(NamedTuple):
|
||||
hidden_states: torch.Tensor
|
||||
hidden_states_scale: Optional[torch.Tensor]
|
||||
topk_output: TopKOutput
|
||||
# SGLANG_OPT_MOE_QUANT_ONCE: optional pre-quantized (q, scale) pair for
|
||||
# ``hidden_states`` (per-token-group-128 fp8, q rows possibly padded to a
|
||||
# multiple of 4). Consumed by the standard->triton fused runner so it can
|
||||
# skip its own activation quant; ``hidden_states`` itself stays bf16.
|
||||
hidden_states_pre_quant: Optional[Tuple[torch.Tensor, torch.Tensor]] = None
|
||||
|
||||
@property
|
||||
def format(self) -> DispatchOutputFormat:
|
||||
@@ -213,9 +224,18 @@ class StandardDispatcher(BaseDispatcher):
|
||||
)
|
||||
elif not self.use_aiter_moe_runner:
|
||||
if TopKOutputChecker.format_is_standard(topk_output):
|
||||
topk_output = topk_output._replace(
|
||||
topk_ids=self.local_expert_mapping[topk_output.topk_ids]
|
||||
)
|
||||
topk_ids_local = self.local_expert_mapping[topk_output.topk_ids]
|
||||
# Drop dp-attention MAX_LEN pad rows from the dispatch:
|
||||
# pad rows carry stale hidden through the router and
|
||||
# their expert outputs are discarded downstream — pure
|
||||
# wasted compute (and a masked-grouped-GEMM workspace
|
||||
# blow-up when they collide on the same top-k). Must
|
||||
# run POST-translation (a pre-translation -1 aliases to
|
||||
# the mapping table's last entry); -1 is the drop
|
||||
# sentinel both the triton and deep_gemm runners honor.
|
||||
if _MASK_DP_PAD_MOE and is_dp_max_padding():
|
||||
mask_dp_pad_moe_topk_ids(topk_ids_local)
|
||||
topk_output = topk_output._replace(topk_ids=topk_ids_local)
|
||||
elif TopKOutputChecker.format_is_triton_kernels(topk_output):
|
||||
raise NotImplementedError()
|
||||
|
||||
|
||||
@@ -789,11 +789,28 @@ def cutlass_w8a8_block_fp8_linear_with_fallback(
|
||||
input_scale: Optional[torch.Tensor] = None,
|
||||
bias: Optional[torch.Tensor] = None,
|
||||
) -> torch.Tensor:
|
||||
assert input_scale is None
|
||||
|
||||
# TODO: add more robust shape check here
|
||||
shape_supported = weight.shape[0] % 128 == 0 and weight.shape[1] % 128 == 0
|
||||
|
||||
if input_scale is not None:
|
||||
# Pre-quantized activation (SGLANG_OPT_MOE_QUANT_ONCE): ``input`` is
|
||||
# the fp8 per-token-group-128 q (rows possibly padded to a multiple
|
||||
# of 4), ``input_scale`` the matching column-major scales
|
||||
# (stride(0) == 1). Output keeps the (padded) row count; the caller
|
||||
# slices back to the true token count.
|
||||
assert shape_supported, (
|
||||
"pre-quantized fp8 input requires cutlass-supported weight shapes "
|
||||
f"(got {tuple(weight.shape)})"
|
||||
)
|
||||
assert input.dtype == torch.float8_e4m3fn
|
||||
input_2d = input.view(-1, input.shape[-1])
|
||||
output = fp8_blockwise_scaled_mm(
|
||||
input_2d, weight.T, input_scale, weight_scale.T, out_dtype=torch.bfloat16
|
||||
)
|
||||
if bias is not None:
|
||||
output += bias
|
||||
return output.view(*input.shape[:-1], weight.shape[0])
|
||||
|
||||
if not shape_supported:
|
||||
# fallback to triton
|
||||
return triton_w8a8_block_fp8_linear(
|
||||
@@ -829,7 +846,33 @@ def deepgemm_w8a8_block_fp8_linear_with_fallback(
|
||||
input_scale: Optional[torch.Tensor] = None,
|
||||
bias: Optional[torch.Tensor] = None,
|
||||
) -> torch.Tensor:
|
||||
assert input_scale is None
|
||||
if input_scale is not None:
|
||||
# Pre-quantized activation (SGLANG_OPT_MOE_QUANT_ONCE): ``input`` is
|
||||
# the fp8 per-token-group-128 q with rows padded to a multiple of 4
|
||||
# and ``input_scale`` the matching column-major fp32 scales
|
||||
# (stride == (1, padded_rows)) -- identical to the MN-major
|
||||
# TMA-aligned layout this path's own quant would produce below.
|
||||
# Output keeps the padded row count; the caller slices back.
|
||||
# UE8M0 packed scales (Blackwell DeepGEMM) use a different layout;
|
||||
# the caller gates on it.
|
||||
assert not deep_gemm_wrapper.DEEPGEMM_SCALE_UE8M0
|
||||
assert input.dtype == torch.float8_e4m3fn
|
||||
assert weight.shape[0] % 64 == 0 and weight.shape[1] % 128 == 0, (
|
||||
"pre-quantized fp8 input requires DeepGEMM-supported weight shapes "
|
||||
f"(got {tuple(weight.shape)})"
|
||||
)
|
||||
input_2d = input.view(-1, input.shape[-1])
|
||||
output = w8a8_block_fp8_matmul_deepgemm(
|
||||
input_2d,
|
||||
weight,
|
||||
input_scale,
|
||||
weight_scale,
|
||||
block_size,
|
||||
output_dtype=torch.bfloat16,
|
||||
)
|
||||
if bias is not None:
|
||||
output += bias
|
||||
return output.view(*input.shape[:-1], weight.shape[0])
|
||||
|
||||
output_dtype = input.dtype
|
||||
dtype_supported = output_dtype == torch.bfloat16
|
||||
|
||||
@@ -1282,6 +1282,7 @@ class ForwardBatch(ForwardBatchDeepSeekMHAMixin):
|
||||
num_tokens,
|
||||
dp_padding_mode.is_max_len(),
|
||||
global_num_tokens,
|
||||
self.global_num_tokens_gpu,
|
||||
)
|
||||
set_is_extend_in_batch(self.is_extend_in_batch)
|
||||
|
||||
|
||||
@@ -6,6 +6,7 @@ from typing import TYPE_CHECKING, Optional
|
||||
|
||||
import torch
|
||||
|
||||
from sglang.kernels.ops.kvcache.cache_ops import absorbed_bmm_concat_cast_q_fp8
|
||||
from sglang.kernels.ops.quantization.fp8_kernel import (
|
||||
fp8_dtype,
|
||||
per_tensor_quant_mla_fp8,
|
||||
@@ -52,6 +53,7 @@ from sglang.srt.model_executor.runner_backend_utils.breakable_cuda_graph.context
|
||||
is_in_breakable_cuda_graph,
|
||||
)
|
||||
from sglang.srt.model_executor.runner_backend_utils.tc_piecewise_cuda_graph import (
|
||||
get_tc_piecewise_forward_context,
|
||||
is_in_tc_piecewise_cuda_graph,
|
||||
)
|
||||
from sglang.srt.models.deepseek_common.utils import (
|
||||
@@ -75,6 +77,8 @@ from sglang.srt.utils.custom_op import register_custom_op
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
_SGLANG_EXPERIMENTAL_LORA_OPTI = envs.SGLANG_EXPERIMENTAL_LORA_OPTI.get()
|
||||
_ENABLE_DSA_Q8KV8_BORN_FP8_Q = envs.SGLANG_ENABLE_DSA_Q8KV8_BORN_FP8_Q.get()
|
||||
_ENABLE_DSA_Q8KV8_QPREP_OVERLAP = envs.SGLANG_ENABLE_DSA_Q8KV8_QPREP_OVERLAP.get()
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from sglang.srt.models.deepseek_v2 import DeepseekV2AttentionMLA
|
||||
@@ -254,6 +258,84 @@ class DeepseekMLAForwardMixin:
|
||||
attn_output_buf=attn_output_buf,
|
||||
)
|
||||
|
||||
def _q8kv8_born_fp8_q_backend(
|
||||
self: DeepseekV2AttentionMLA,
|
||||
forward_batch: ForwardBatch,
|
||||
llama_4_scaling: Optional[torch.Tensor],
|
||||
):
|
||||
"""Return the DSA backend iff the born-fp8 q fast path can run.
|
||||
|
||||
Gated by SGLANG_ENABLE_DSA_Q8KV8_BORN_FP8_Q (checked by the caller).
|
||||
When this returns a backend, the bf16 absorbed bmm + the standalone
|
||||
concat_and_cast_q_fp8_pad are replaced by one fused kernel that writes
|
||||
the fp8 q directly into the backend's q8kv8 buffer; q_nope_out becomes
|
||||
a NaN sentinel. Every condition here must therefore guarantee that
|
||||
forward_extend consumes q via _forward_flashmla_sparse_q8kv8 and that
|
||||
nothing else reads q_nope_out's payload.
|
||||
"""
|
||||
from sglang.srt.model_executor.runner import get_is_capture_mode
|
||||
|
||||
if llama_4_scaling is not None:
|
||||
return None
|
||||
if _is_hip or _is_cpu:
|
||||
return None
|
||||
if self.current_attention_backend not in FORWARD_ABSORB_CORE_ATTENTION_BACKENDS:
|
||||
return None
|
||||
if self.use_deep_gemm_bmm:
|
||||
return None
|
||||
w_kc = self.w_kc
|
||||
if w_kc is None or w_kc.dtype != torch.bfloat16:
|
||||
return None
|
||||
if is_kv_b_lora_active(self) or _SGLANG_EXPERIMENTAL_LORA_OPTI:
|
||||
return None
|
||||
# The fused kernel consumes the post-rope q_pe, so the eager rope
|
||||
# apply below must run (mirror of its condition).
|
||||
if self.rotary_emb is None:
|
||||
return None
|
||||
if self._fuse_rope_for_trtllm_mla(forward_batch):
|
||||
return None
|
||||
if self._skip_rope_for_dsa_tilelang_fused():
|
||||
return None
|
||||
if self._skip_rope_for_aiter_fused_mla():
|
||||
return None
|
||||
if _use_aiter and _is_gfx95_supported and not self.use_dsa:
|
||||
return None
|
||||
# Graph/compile surfaces run their own dispatch; the python-side
|
||||
# stash handshake is eager-only.
|
||||
if is_graph_dsa_split_op_surface(forward_batch):
|
||||
return None
|
||||
if get_tc_piecewise_forward_context() is not None:
|
||||
return None
|
||||
if is_in_breakable_cuda_graph():
|
||||
return None
|
||||
if get_is_capture_mode():
|
||||
return None
|
||||
if get_parallel().dcp_enabled:
|
||||
return None
|
||||
# Context-parallel prefill reshuffles the KV side; keep the handshake
|
||||
# out of those paths.
|
||||
if dsa_use_prefill_cp(forward_batch) or mla_use_prefill_cp(forward_batch):
|
||||
return None
|
||||
# Kernel shape constraints (tl.arange / tl.dot / block tiling). K
|
||||
# (qk_nope_head_dim) needs only K % 16 == 0 and K <= 256: power-of-2
|
||||
# K (DeepSeek 128) takes the kernel's preload-once path, other K
|
||||
# (GLM-5 192) its split-K loop.
|
||||
k_dim = self.qk_nope_head_dim
|
||||
rope_dim = self.qk_rope_head_dim
|
||||
if k_dim < 16 or k_dim > 256 or k_dim % 16 != 0:
|
||||
return None
|
||||
if rope_dim <= 0 or (rope_dim & (rope_dim - 1)) != 0:
|
||||
return None
|
||||
if self.kv_lora_rank % 128 != 0:
|
||||
return None
|
||||
if tuple(w_kc.shape) != (self.num_local_heads, k_dim, self.kv_lora_rank):
|
||||
return None
|
||||
backend = get_attn_backend()
|
||||
eligible = getattr(backend, "q8kv8_born_fp8_q_eligible", None)
|
||||
if eligible is None or not eligible(forward_batch, self.num_local_heads):
|
||||
return None
|
||||
return backend
|
||||
|
||||
def forward_absorb_prepare(
|
||||
self: DeepseekV2AttentionMLA,
|
||||
positions: torch.Tensor,
|
||||
@@ -265,6 +347,11 @@ class DeepseekMLAForwardMixin:
|
||||
):
|
||||
from sglang.srt.model_executor.runner import get_is_capture_mode
|
||||
|
||||
# Q8KV8 q-prep/indexer overlap handshake (see the fork site below):
|
||||
# True between the alt-stream fork and its consumption in the born
|
||||
# block; also suppresses the duplicate split/rope on that path.
|
||||
self._q8kv8_qprep_overlap_pending = False
|
||||
|
||||
fuse_bmm_attention = (
|
||||
self.q_lora_rank is not None
|
||||
and self._can_fuse_bmm_into_attention(forward_batch)
|
||||
@@ -422,6 +509,47 @@ class DeepseekMLAForwardMixin:
|
||||
q_nope, q_pe, k_pe = self._split_q_nope_pe(q, latent_cache)
|
||||
fusion_plan = self._make_mla_bmm_fusion_plan(q, q_nope)
|
||||
|
||||
# Q8KV8 q-prep/indexer overlap (opt-in): the born-fp8 q-prep
|
||||
# chain (split -> rope -> fused absorbed-bmm+cast, ~173us)
|
||||
# and the indexer chain both fork from the q_a_layernorm
|
||||
# output and never touch each other's tensors, so the q-prep
|
||||
# can run on alt_stream underneath the indexer. The fork
|
||||
# must be enqueued BEFORE the indexer (a later wait_stream
|
||||
# would serialize behind it). The born predicate itself
|
||||
# guarantees eager-only and the plain-rope branch (all fused
|
||||
# /skip-rope variants make it return None), so applying rope
|
||||
# here is exactly what the skipped block below would do.
|
||||
if (
|
||||
_ENABLE_DSA_Q8KV8_QPREP_OVERLAP
|
||||
and _ENABLE_DSA_Q8KV8_BORN_FP8_Q
|
||||
and fusion_plan is None
|
||||
and self.alt_stream is not None
|
||||
and q_lora is not None
|
||||
and self.rotary_emb is not None
|
||||
):
|
||||
_born_backend_early = self._q8kv8_born_fp8_q_backend(
|
||||
forward_batch, llama_4_scaling
|
||||
)
|
||||
if _born_backend_early is not None:
|
||||
q_nope, q_pe, k_pe = self._split_q_nope_pe(q, latent_cache)
|
||||
q_pe, k_pe = self.rotary_emb(positions, q_pe, k_pe)
|
||||
_q_fp8 = _born_backend_early.q8kv8_acquire_born_q_buffer(
|
||||
q_nope.shape[0],
|
||||
self.num_local_heads,
|
||||
self.kv_lora_rank + self.qk_rope_head_dim,
|
||||
q_nope.device,
|
||||
)
|
||||
self.alt_stream.wait_stream(torch.cuda.current_stream())
|
||||
with torch.cuda.stream(self.alt_stream):
|
||||
absorbed_bmm_concat_cast_q_fp8(
|
||||
_q_fp8,
|
||||
q_nope,
|
||||
self.w_kc,
|
||||
q_pe,
|
||||
self.num_local_heads,
|
||||
)
|
||||
self._q8kv8_qprep_overlap_pending = True
|
||||
|
||||
if q_lora is not None:
|
||||
if self.should_run_indexer(prev_topk_indices):
|
||||
topk_indices = self.indexer(
|
||||
@@ -456,6 +584,15 @@ class DeepseekMLAForwardMixin:
|
||||
q_nope, q_pe, k_pe = self._split_q_nope_pe(q, latent_cache)
|
||||
|
||||
_kvb_q = None
|
||||
born_q_backend = None
|
||||
if (
|
||||
_ENABLE_DSA_Q8KV8_BORN_FP8_Q
|
||||
and fusion_plan is None
|
||||
and q_nope.dtype == torch.bfloat16
|
||||
):
|
||||
born_q_backend = self._q8kv8_born_fp8_q_backend(
|
||||
forward_batch, llama_4_scaling
|
||||
)
|
||||
if q_replicate_active:
|
||||
# full-head absorb with the pre-gathered w_kc (q_nope already full-head)
|
||||
q_nope_out = (
|
||||
@@ -467,6 +604,11 @@ class DeepseekMLAForwardMixin:
|
||||
# The composite split op fills q_nope_out_buf and attention reads
|
||||
# this transposed alias directly.
|
||||
q_nope_out = fusion_plan.q_nope_out_view
|
||||
elif born_q_backend is not None:
|
||||
# Born-fp8 q: skip the bf16 absorbed bmm entirely; the fused
|
||||
# bmm+concat+cast kernel (launched after rope below) writes the
|
||||
# fp8 q directly into the q8kv8 backend buffer.
|
||||
q_nope_out = None
|
||||
else:
|
||||
if _SGLANG_EXPERIMENTAL_LORA_OPTI:
|
||||
# Fork the kv_b q-correction A-step onto the LoRA side stream to overlap the bmm.
|
||||
@@ -591,9 +733,41 @@ class DeepseekMLAForwardMixin:
|
||||
or self.use_dsa
|
||||
or self.current_attention_backend == "triton"
|
||||
)
|
||||
# Already applied at the q-prep/indexer overlap fork.
|
||||
and not self._q8kv8_qprep_overlap_pending
|
||||
):
|
||||
q_pe, k_pe = self.rotary_emb(positions, q_pe, k_pe)
|
||||
|
||||
if born_q_backend is not None:
|
||||
# Born-fp8 q (SGLANG_ENABLE_DSA_Q8KV8_BORN_FP8_Q): one fused
|
||||
# kernel replaces bmm -> bf16 q_nope_out ->
|
||||
# concat_and_cast_q_fp8_pad. q_nope is the pre-absorb bf16 view
|
||||
# (rope only touched the disjoint q_pe columns) and q_pe carries
|
||||
# the post-rope values. The stash is consumed by
|
||||
# _forward_flashmla_sparse_q8kv8; q_nope_out becomes a
|
||||
# NaN-poisoned shape-only sentinel.
|
||||
num_tokens = q_nope.shape[0]
|
||||
if self._q8kv8_qprep_overlap_pending:
|
||||
# q_fp8 was produced on alt_stream at the fork above; join so
|
||||
# everything downstream (incl. the next layer's fork, which
|
||||
# reuses the single born-q slot) orders after it.
|
||||
torch.cuda.current_stream().wait_stream(self.alt_stream)
|
||||
self._q8kv8_qprep_overlap_pending = False
|
||||
else:
|
||||
q_fp8 = born_q_backend.q8kv8_acquire_born_q_buffer(
|
||||
num_tokens,
|
||||
self.num_local_heads,
|
||||
self.kv_lora_rank + self.qk_rope_head_dim,
|
||||
q_nope.device,
|
||||
)
|
||||
absorbed_bmm_concat_cast_q_fp8(
|
||||
q_fp8, q_nope, self.w_kc, q_pe, self.num_local_heads
|
||||
)
|
||||
born_q_backend.q8kv8_stash_born_q(num_tokens, self.attn_mqa.layer_id)
|
||||
q_nope_out = born_q_backend.q8kv8_born_q_sentinel(
|
||||
num_tokens, self.num_local_heads, self.kv_lora_rank, q_nope.device
|
||||
)
|
||||
|
||||
dsa_prefill_cp = dsa_use_prefill_cp(forward_batch)
|
||||
mla_prefill_cp = mla_use_prefill_cp(forward_batch)
|
||||
defer_kv_gather_until_after_rope = _should_defer_dsa_cp_kv_gather(
|
||||
|
||||
@@ -241,6 +241,9 @@ from sglang.kernels.ops.gemm.fused_a_gemm import (
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
# One-time SGLANG_OPT_MOE_QUANT_ONCE engagement log (see _moe_quant_once_enabled).
|
||||
_moe_quant_once_logged = False
|
||||
|
||||
_enable_pcg_dsv2_dual_stream = (
|
||||
_is_cuda and envs.SGLANG_ENABLE_PCG_DSV2_DUAL_STREAM.get()
|
||||
)
|
||||
@@ -307,6 +310,7 @@ class DeepseekV2MLP(nn.Module):
|
||||
x,
|
||||
forward_batch=None,
|
||||
gemm_output_zero_allocator: BumpAllocator = None,
|
||||
gateup_pre_quant: Optional[Tuple[torch.Tensor, torch.Tensor]] = None,
|
||||
):
|
||||
if (self.tp_size == 1) and x.shape[0] == 0:
|
||||
return x
|
||||
@@ -336,17 +340,24 @@ class DeepseekV2MLP(nn.Module):
|
||||
out, _ = self.down_proj((out_fp4, out_scale))
|
||||
return out
|
||||
|
||||
if (
|
||||
gemm_output_zero_allocator is not None
|
||||
and x.shape[0] <= 256
|
||||
and self.gate_up_proj.weight.dtype == torch.uint8
|
||||
):
|
||||
y = gemm_output_zero_allocator.allocate(
|
||||
x.shape[0] * self.gate_up_proj.output_size_per_partition
|
||||
).view(x.shape[0], self.gate_up_proj.output_size_per_partition)
|
||||
x = (x, None, y)
|
||||
if gateup_pre_quant is not None:
|
||||
# SGLANG_OPT_MOE_QUANT_ONCE: reuse the caller's per-token-group-128
|
||||
# fp8 (q, scale) of x for the gate_up GEMM instead of re-quantizing
|
||||
# inside the fp8 linear method. q rows may be padded to a multiple
|
||||
# of 4; the caller slices the MLP output back.
|
||||
gate_up, _ = self.gate_up_proj(gateup_pre_quant)
|
||||
else:
|
||||
if (
|
||||
gemm_output_zero_allocator is not None
|
||||
and x.shape[0] <= 256
|
||||
and self.gate_up_proj.weight.dtype == torch.uint8
|
||||
):
|
||||
y = gemm_output_zero_allocator.allocate(
|
||||
x.shape[0] * self.gate_up_proj.output_size_per_partition
|
||||
).view(x.shape[0], self.gate_up_proj.output_size_per_partition)
|
||||
x = (x, None, y)
|
||||
|
||||
gate_up, _ = self.gate_up_proj(x)
|
||||
gate_up, _ = self.gate_up_proj(x)
|
||||
# Fast path: fused silu+clamp+fp8_quant+deepgemm when conditions met.
|
||||
# Only valid when down_proj does NOT need an all-reduce and its weights
|
||||
# are fp8 (uint8 storage with weight_scale_inv).
|
||||
@@ -822,6 +833,9 @@ class DeepseekV2MoE(nn.Module):
|
||||
or get_moe_a2a_backend().is_flashinfer()
|
||||
)
|
||||
self._fuse_shared_experts_inside_sbo = SboFlags.fuse_shared_experts_inside_sbo()
|
||||
# SGLANG_OPT_MOE_QUANT_ONCE eligibility, resolved lazily on first
|
||||
# forward (weights and runner are final by then). None = undecided.
|
||||
self._moe_quant_once: Optional[bool] = None
|
||||
|
||||
def get_moe_weights(self):
|
||||
# EPLB only rebalances physical routed experts. Fused shared expert
|
||||
@@ -933,6 +947,13 @@ class DeepseekV2MoE(nn.Module):
|
||||
# deep_gemm does not free hidden_states, which the shared expert reads on the alt stream.
|
||||
use_flashinfer_trtllm_bypass = get_forward().flashinfer_trtllm_bypass
|
||||
current_stream = torch.cuda.current_stream()
|
||||
# Quantize-once (SGLANG_OPT_MOE_QUANT_ONCE) must happen on the main
|
||||
# stream BEFORE the alt-stream fork so both consumers see it.
|
||||
pre_quant_input = (
|
||||
None
|
||||
if use_flashinfer_trtllm_bypass
|
||||
else self._maybe_quant_moe_input_once(hidden_states)
|
||||
)
|
||||
self.alt_stream.wait_stream(current_stream)
|
||||
has_shared_output = (
|
||||
hidden_states.shape[0] > 0 and self.num_fused_shared_experts == 0
|
||||
@@ -975,6 +996,10 @@ class DeepseekV2MoE(nn.Module):
|
||||
)
|
||||
elif use_flashinfer_trtllm_bypass:
|
||||
final_hidden_states = self.experts.forward_impl(hidden_states, topk_output)
|
||||
elif pre_quant_input is not None:
|
||||
final_hidden_states = self.experts(
|
||||
hidden_states, topk_output, pre_quant_input=pre_quant_input
|
||||
)
|
||||
else:
|
||||
final_hidden_states = self.experts(hidden_states, topk_output)
|
||||
if (
|
||||
@@ -988,7 +1013,9 @@ class DeepseekV2MoE(nn.Module):
|
||||
# Shared expert on alt stream, issued AFTER the main (routed) branch. See note above.
|
||||
with torch.cuda.stream(self.alt_stream):
|
||||
shared_output = self._forward_shared_experts(
|
||||
hidden_states, gemm_output_zero_allocator
|
||||
hidden_states,
|
||||
gemm_output_zero_allocator,
|
||||
pre_quant_input=pre_quant_input,
|
||||
)
|
||||
|
||||
current_stream.wait_stream(self.alt_stream)
|
||||
@@ -1044,13 +1071,22 @@ class DeepseekV2MoE(nn.Module):
|
||||
# reduce_scatterv. When set, never compute/add it here (on the global buffer).
|
||||
shared_output = None
|
||||
if hidden_states.shape[0] > 0:
|
||||
# Quantize-once (SGLANG_OPT_MOE_QUANT_ONCE): only worthwhile when
|
||||
# the shared expert also runs here on the same tensor.
|
||||
pre_quant_input = (
|
||||
None
|
||||
if skip_shared_experts
|
||||
else self._maybe_quant_moe_input_once(hidden_states)
|
||||
)
|
||||
if (
|
||||
not defer_shared
|
||||
and not self._fuse_shared_experts_inside_sbo
|
||||
and not skip_shared_experts
|
||||
):
|
||||
shared_output = self._forward_shared_experts(
|
||||
hidden_states, gemm_output_zero_allocator
|
||||
hidden_states,
|
||||
gemm_output_zero_allocator,
|
||||
pre_quant_input=pre_quant_input,
|
||||
)
|
||||
# router_logits: (num_tokens, n_experts)
|
||||
router_logits = self.gate(hidden_states, gemm_output_zero_allocator)
|
||||
@@ -1066,6 +1102,7 @@ class DeepseekV2MoE(nn.Module):
|
||||
**topk_kwargs,
|
||||
)
|
||||
else:
|
||||
pre_quant_input = None
|
||||
shared_output = None
|
||||
topk_output = self.topk.empty_topk_output(
|
||||
hidden_states.device, layer_id=self.layer_id
|
||||
@@ -1101,10 +1138,17 @@ class DeepseekV2MoE(nn.Module):
|
||||
self.experts.dispatcher.register_post_combine_hook(_post_combine_hook)
|
||||
)
|
||||
|
||||
final_hidden_states = self.experts(
|
||||
hidden_states,
|
||||
topk_output,
|
||||
)
|
||||
if pre_quant_input is not None:
|
||||
final_hidden_states = self.experts(
|
||||
hidden_states,
|
||||
topk_output,
|
||||
pre_quant_input=pre_quant_input,
|
||||
)
|
||||
else:
|
||||
final_hidden_states = self.experts(
|
||||
hidden_states,
|
||||
topk_output,
|
||||
)
|
||||
if (
|
||||
not _is_cuda
|
||||
and not _is_musa
|
||||
@@ -1122,7 +1166,9 @@ class DeepseekV2MoE(nn.Module):
|
||||
and not skip_shared_experts
|
||||
):
|
||||
shared_output = self._forward_shared_experts(
|
||||
hidden_states, gemm_output_zero_allocator
|
||||
hidden_states,
|
||||
gemm_output_zero_allocator,
|
||||
pre_quant_input=pre_quant_input,
|
||||
)
|
||||
|
||||
final_hidden_states = maybe_fuse_routed_scale_and_shared_add(
|
||||
@@ -1426,15 +1472,131 @@ class DeepseekV2MoE(nn.Module):
|
||||
return final_hidden_states
|
||||
|
||||
def _forward_shared_experts(
|
||||
self, hidden_states, gemm_output_zero_allocator: BumpAllocator = None
|
||||
self,
|
||||
hidden_states,
|
||||
gemm_output_zero_allocator: BumpAllocator = None,
|
||||
pre_quant_input: Optional[Tuple[torch.Tensor, torch.Tensor]] = None,
|
||||
):
|
||||
if (hidden_states.shape[0] > 0) and (self.num_fused_shared_experts == 0):
|
||||
if pre_quant_input is not None:
|
||||
# SGLANG_OPT_MOE_QUANT_ONCE: (q, s) rows may be padded to a
|
||||
# multiple of 4; the padded rows flow through the MLP (all ops
|
||||
# are row-local) and are sliced off here.
|
||||
out = self.shared_experts(
|
||||
hidden_states, gateup_pre_quant=pre_quant_input
|
||||
)
|
||||
return out[: hidden_states.shape[0]]
|
||||
return self.shared_experts(
|
||||
hidden_states, gemm_output_zero_allocator=gemm_output_zero_allocator
|
||||
)
|
||||
else:
|
||||
return None
|
||||
|
||||
def _moe_quant_once_enabled(self) -> bool:
|
||||
"""SGLANG_OPT_MOE_QUANT_ONCE: quantize the (dp-gathered) MoE input to
|
||||
per-token-group-128 fp8 once per layer and feed both the fused shared
|
||||
expert's fp8 GEMM (cutlass or deepgemm w8a8 linear) and the routed
|
||||
experts' triton fused runner, instead of quantizing the same
|
||||
[T, hidden] tensor twice with different scale layouts."""
|
||||
if self._moe_quant_once is None:
|
||||
self._moe_quant_once, reason = self._compute_moe_quant_once_enabled()
|
||||
global _moe_quant_once_logged
|
||||
if envs.SGLANG_OPT_MOE_QUANT_ONCE.get() and not _moe_quant_once_logged:
|
||||
_moe_quant_once_logged = True
|
||||
logger.info(
|
||||
"SGLANG_OPT_MOE_QUANT_ONCE: %s (layer %s)",
|
||||
"ENGAGED" if self._moe_quant_once else f"INELIGIBLE: {reason}",
|
||||
self.layer_id,
|
||||
)
|
||||
return self._moe_quant_once
|
||||
|
||||
def _compute_moe_quant_once_enabled(self) -> Tuple[bool, str]:
|
||||
"""Returns (eligible, reason); reason names the first failing check."""
|
||||
from sglang.srt.layers.moe.token_dispatcher.standard import StandardDispatcher
|
||||
from sglang.srt.layers.quantization.fp8 import Fp8LinearMethod, Fp8MoEMethod
|
||||
from sglang.srt.layers.quantization.fp8_utils import (
|
||||
cutlass_w8a8_block_fp8_linear_with_fallback,
|
||||
deepgemm_w8a8_block_fp8_linear_with_fallback,
|
||||
)
|
||||
|
||||
if not envs.SGLANG_OPT_MOE_QUANT_ONCE.get():
|
||||
return False, "env off"
|
||||
if not _is_cuda:
|
||||
return False, "not CUDA"
|
||||
if self._enable_a2a_moe or self._fuse_shared_experts_inside_sbo:
|
||||
return False, "a2a MoE or SBO shared-expert fusion"
|
||||
# Shared-expert side: fp8 block-128 weights served by a w8a8 linear
|
||||
# backend taught to accept a pre-quantized (q, scale) tuple: cutlass
|
||||
# or deepgemm (fp32 scales only, i.e. not UE8M0/Blackwell).
|
||||
if self.num_fused_shared_experts != 0 or not hasattr(self, "shared_experts"):
|
||||
return False, "no separate shared experts"
|
||||
if not self.shared_experts_is_fp8:
|
||||
return False, "shared experts not fp8"
|
||||
if self.shared_experts_weight_block_size != [128, 128]:
|
||||
return False, "shared weight block size != [128, 128]"
|
||||
gate_up = self.shared_experts.gate_up_proj
|
||||
if not isinstance(gate_up.quant_method, Fp8LinearMethod):
|
||||
return False, "shared gate_up quant method not Fp8LinearMethod"
|
||||
linear_fn = gate_up.quant_method.w8a8_block_fp8_linear
|
||||
if linear_fn is cutlass_w8a8_block_fp8_linear_with_fallback:
|
||||
if gate_up.weight.shape[0] % 128 != 0 or gate_up.weight.shape[1] % 128 != 0:
|
||||
return False, "gate_up weight shape unsupported by cutlass"
|
||||
elif linear_fn is deepgemm_w8a8_block_fp8_linear_with_fallback:
|
||||
from sglang.srt.layers import deep_gemm_wrapper
|
||||
|
||||
if deep_gemm_wrapper.DEEPGEMM_SCALE_UE8M0:
|
||||
return False, "DeepGEMM UE8M0 scales (Blackwell) unsupported"
|
||||
if gate_up.weight.shape[0] % 64 != 0 or gate_up.weight.shape[1] % 128 != 0:
|
||||
return False, "gate_up weight shape unsupported by deepgemm"
|
||||
else:
|
||||
return False, f"w8a8 linear backend {linear_fn.__name__} unsupported"
|
||||
# Routed side: standard dispatcher + triton fused func with dynamic
|
||||
# per-token-group-128 fp8 activation quant.
|
||||
experts = self.experts
|
||||
if not isinstance(experts, FusedMoE):
|
||||
return False, "experts not FusedMoE"
|
||||
quant_method = experts.quant_method
|
||||
if not isinstance(quant_method, Fp8MoEMethod):
|
||||
return False, "experts quant method not Fp8MoEMethod"
|
||||
if not quant_method.block_quant or quant_method.use_mxfp8:
|
||||
return False, "experts not block-quant fp8"
|
||||
if quant_method.quant_config.weight_block_size != [128, 128]:
|
||||
return False, "experts weight block size != [128, 128]"
|
||||
# Fp8MoEMethod only sets .runner for runner backends it drives itself.
|
||||
runner = getattr(quant_method, "runner", None)
|
||||
if runner is None or not runner.runner_backend.is_triton():
|
||||
return False, "MoE runner backend not triton"
|
||||
if runner.fused_func is None or runner.lora_enabled:
|
||||
return False, "triton fused func unavailable (or LoRA enabled)"
|
||||
if not isinstance(experts.dispatcher, StandardDispatcher):
|
||||
return False, "dispatcher not StandardDispatcher"
|
||||
if experts.moe_runner_config.apply_router_weight_on_input:
|
||||
return False, "apply_router_weight_on_input"
|
||||
if experts.w13_input_scale is not None:
|
||||
return False, "static w13 input scale"
|
||||
return True, "ok"
|
||||
|
||||
def _maybe_quant_moe_input_once(
|
||||
self, hidden_states: torch.Tensor
|
||||
) -> Optional[Tuple[torch.Tensor, torch.Tensor]]:
|
||||
"""Quantize hidden_states once (per-token-group-128 fp8, rows padded to
|
||||
a multiple of 4, column-major scales) for both the shared-expert GEMM
|
||||
and the routed dispatch, or return None when ineligible."""
|
||||
if hidden_states.shape[0] == 0 or hidden_states.dtype != torch.bfloat16:
|
||||
return None
|
||||
if not self._moe_quant_once_enabled():
|
||||
return None
|
||||
if is_in_tc_piecewise_cuda_graph():
|
||||
# The piecewise MoE op quantizes internally; a pre-quant here
|
||||
# would be dead work.
|
||||
return None
|
||||
from sglang.kernels.ops.quantization.fp8_kernel import (
|
||||
sglang_per_token_group_quant_fp8_row_padded,
|
||||
)
|
||||
|
||||
q, s = sglang_per_token_group_quant_fp8_row_padded(hidden_states, 128)
|
||||
return q, s
|
||||
|
||||
def op_gate(self, state):
|
||||
if state.hidden_states_mlp_input.shape[0] > 0:
|
||||
# router_logits: (num_tokens, n_experts)
|
||||
|
||||
@@ -0,0 +1,274 @@
|
||||
"""Standalone GPU test for SGLANG_OPT_MOE_QUANT_ONCE (quantize the MoE input
|
||||
once, feed both the fused shared-expert GEMM and the routed triton runner).
|
||||
|
||||
|
||||
CUDA_VISIBLE_DEVICES=0 python test/manual/test_moe_quant_once.py
|
||||
|
||||
Verifies, against the double-quant baseline:
|
||||
(1) quant equivalence: the row-padded quantize-once kernel produces the
|
||||
same q bits / scale values as the routed path's default row-major quant
|
||||
(JIT v2 kernel) on the valid rows;
|
||||
(2) shared consumer: cutlass_w8a8_block_fp8_linear_with_fallback with a
|
||||
pre-quantized (q, s) tuple vs its own internal quant -- expected BITWISE
|
||||
(baseline uses the identical row-padded quant + identical GEMM);
|
||||
(2b) shared consumer under SGLANG_ENABLE_JIT_DEEPGEMM=1 (the recommended JIT-DeepGEMM config):
|
||||
deepgemm_w8a8_block_fp8_linear_with_fallback with the same (q, s) tuple
|
||||
-- expected BITWISE (DG's own quant layout, column-major TMA-aligned
|
||||
fp32 scales, is byte-identical to the row-padded quantize-once layout);
|
||||
skipped cleanly when deep_gemm is unavailable or UE8M0 (Blackwell);
|
||||
(3) routed consumer: fused_experts(a1_q=..., a1_scale=...) vs the in-kernel
|
||||
quant baseline -- expected BITWISE if (1) is bitwise (the fused kernel
|
||||
reads A_scale through explicit strides, so the column-major scale view
|
||||
feeds identical values).
|
||||
|
||||
If (1) is not bitwise (AOT v2 vs JIT v2 quant kernels round differently),
|
||||
(3) falls back to an allclose check at atol=1e-2 and the discrepancy is
|
||||
reported --.
|
||||
"""
|
||||
|
||||
import sys
|
||||
|
||||
import torch
|
||||
|
||||
from sglang.kernels.ops.quantization.fp8_kernel import (
|
||||
sglang_per_token_group_quant_fp8,
|
||||
sglang_per_token_group_quant_fp8_row_padded,
|
||||
)
|
||||
from sglang.srt.layers.moe.moe_runner.base import MoeRunnerConfig
|
||||
from sglang.srt.layers.moe.moe_runner.triton_utils.fused_moe import fused_experts
|
||||
from sglang.srt.layers.moe.topk import StandardTopKOutput
|
||||
from sglang.srt.layers.quantization.fp8_utils import (
|
||||
cutlass_w8a8_block_fp8_linear_with_fallback,
|
||||
)
|
||||
from sglang.srt.server_args import ServerArgs, set_global_server_args_for_scheduler
|
||||
|
||||
GROUP = 128
|
||||
FAILURES = []
|
||||
|
||||
|
||||
def _report(name, ok, detail=""):
|
||||
status = "PASS" if ok else "FAIL"
|
||||
print(f"[{status}] {name} {detail}")
|
||||
if not ok:
|
||||
FAILURES.append(name)
|
||||
|
||||
|
||||
def _quant_weight_blockwise(w_bf16, block=128):
|
||||
"""Per-[128,128]-block fp8 weight quant (reference, fp32 math)."""
|
||||
n, k = w_bf16.shape
|
||||
w = w_bf16.float().view(n // block, block, k // block, block)
|
||||
amax = w.abs().amax(dim=(1, 3), keepdim=True).clamp(min=1e-4)
|
||||
scale = amax / torch.finfo(torch.float8_e4m3fn).max
|
||||
q = (w / scale).clamp(-448, 448).to(torch.float8_e4m3fn)
|
||||
return (
|
||||
q.view(n, k),
|
||||
scale.squeeze(1).squeeze(-1).to(torch.float32), # [n/128, k/128]
|
||||
)
|
||||
|
||||
|
||||
def test_quant_equivalence(T, K, device):
|
||||
x = torch.randn(T, K, device=device, dtype=torch.bfloat16) * 3
|
||||
q_ref, s_ref = sglang_per_token_group_quant_fp8(x, GROUP) # routed baseline
|
||||
q_pad, s_pad = sglang_per_token_group_quant_fp8_row_padded(x, GROUP)
|
||||
|
||||
bitwise_q = torch.equal(q_pad[:T].view(torch.uint8), q_ref.view(torch.uint8))
|
||||
bitwise_s = torch.equal(s_pad[:T].contiguous(), s_ref)
|
||||
_report(
|
||||
f"quant-equivalence T={T} K={K}",
|
||||
bitwise_q and bitwise_s,
|
||||
f"(q bitwise={bitwise_q}, s bitwise={bitwise_s})",
|
||||
)
|
||||
return bitwise_q and bitwise_s
|
||||
|
||||
|
||||
def test_shared_consumer(T, K, N, device):
|
||||
torch.manual_seed(T + K)
|
||||
x = torch.randn(T, K, device=device, dtype=torch.bfloat16)
|
||||
w_bf16 = torch.randn(N, K, device=device, dtype=torch.bfloat16) / K**0.5
|
||||
w, ws = _quant_weight_blockwise(w_bf16)
|
||||
|
||||
ref = cutlass_w8a8_block_fp8_linear_with_fallback(x, w, [128, 128], ws)
|
||||
|
||||
q_pad, s_pad = sglang_per_token_group_quant_fp8_row_padded(x, GROUP)
|
||||
out = cutlass_w8a8_block_fp8_linear_with_fallback(
|
||||
q_pad, w, [128, 128], ws, input_scale=s_pad
|
||||
)[:T]
|
||||
|
||||
bitwise = torch.equal(out, ref)
|
||||
close = torch.allclose(out.float(), ref.float(), atol=1e-2, rtol=1e-2)
|
||||
_report(
|
||||
f"shared-consumer T={T} K={K} N={N}",
|
||||
close,
|
||||
f"(bitwise={bitwise}, max|d|={(out.float() - ref.float()).abs().max().item():.3e})",
|
||||
)
|
||||
return bitwise
|
||||
|
||||
|
||||
def test_shared_consumer_deepgemm(T, K, N, device):
|
||||
"""DG branch (SGLANG_ENABLE_JIT_DEEPGEMM=1 recommended JIT-DeepGEMM config): the shared-expert
|
||||
linear resolves to deepgemm_w8a8_block_fp8_linear_with_fallback. Its own
|
||||
quant (column-major + TMA-aligned fp32 scales) has the same buffer layout
|
||||
as the row-padded quantize-once kernel, so this is expected BITWISE."""
|
||||
from sglang.srt.layers.quantization.fp8_utils import (
|
||||
deepgemm_w8a8_block_fp8_linear_with_fallback,
|
||||
)
|
||||
|
||||
torch.manual_seed(T + K + 1)
|
||||
x = torch.randn(T, K, device=device, dtype=torch.bfloat16)
|
||||
w_bf16 = torch.randn(N, K, device=device, dtype=torch.bfloat16) / K**0.5
|
||||
w, ws = _quant_weight_blockwise_n64(w_bf16)
|
||||
|
||||
ref = deepgemm_w8a8_block_fp8_linear_with_fallback(x, w, [128, 128], ws)
|
||||
|
||||
q_pad, s_pad = sglang_per_token_group_quant_fp8_row_padded(x, GROUP)
|
||||
out = deepgemm_w8a8_block_fp8_linear_with_fallback(
|
||||
q_pad, w, [128, 128], ws, input_scale=s_pad
|
||||
)[:T]
|
||||
|
||||
bitwise = torch.equal(out, ref)
|
||||
close = torch.allclose(out.float(), ref.float(), atol=1e-2, rtol=1e-2)
|
||||
_report(
|
||||
f"shared-consumer-deepgemm T={T} K={K} N={N}",
|
||||
close,
|
||||
f"(bitwise={bitwise}, max|d|={(out.float() - ref.float()).abs().max().item():.3e})",
|
||||
)
|
||||
return bitwise
|
||||
|
||||
|
||||
def _quant_weight_blockwise_n64(w_bf16, block=128):
|
||||
"""Like _quant_weight_blockwise but supports N % 64 == 0 (DeepGEMM's
|
||||
minimum): the last (partial) N-block reuses ceil-division block indexing."""
|
||||
n, k = w_bf16.shape
|
||||
if n % block == 0:
|
||||
return _quant_weight_blockwise(w_bf16, block)
|
||||
import math
|
||||
|
||||
n_blocks = math.ceil(n / block)
|
||||
w = w_bf16.float()
|
||||
q = torch.empty(n, k, device=w.device, dtype=torch.float8_e4m3fn)
|
||||
scale = torch.empty(n_blocks, k // block, device=w.device, dtype=torch.float32)
|
||||
for bn in range(n_blocks):
|
||||
rows = slice(bn * block, min((bn + 1) * block, n))
|
||||
wb = w[rows].view(rows.stop - rows.start, k // block, block)
|
||||
amax = wb.abs().amax(dim=(0, 2)).clamp(min=1e-4)
|
||||
s = amax / torch.finfo(torch.float8_e4m3fn).max
|
||||
q[rows] = (
|
||||
(wb / s[None, :, None]).clamp(-448, 448).to(torch.float8_e4m3fn).view(-1, k)
|
||||
)
|
||||
scale[bn] = s
|
||||
return q, scale
|
||||
|
||||
|
||||
def test_routed_consumer(T, K, E, I, topk, device):
|
||||
torch.manual_seed(T * 7 + K)
|
||||
x = torch.randn(T, K, device=device, dtype=torch.bfloat16)
|
||||
w1 = torch.empty(E, 2 * I, K, device=device, dtype=torch.float8_e4m3fn)
|
||||
w1s = torch.empty(E, 2 * I // 128, K // 128, device=device)
|
||||
w2 = torch.empty(E, K, I, device=device, dtype=torch.float8_e4m3fn)
|
||||
w2s = torch.empty(E, K // 128, I // 128, device=device)
|
||||
for e in range(E):
|
||||
w1[e], w1s[e] = _quant_weight_blockwise(
|
||||
torch.randn(2 * I, K, device=device, dtype=torch.bfloat16) / K**0.5
|
||||
)
|
||||
w2[e], w2s[e] = _quant_weight_blockwise(
|
||||
torch.randn(K, I, device=device, dtype=torch.bfloat16) / I**0.5
|
||||
)
|
||||
|
||||
topk_weights = torch.rand(T, topk, device=device)
|
||||
topk_weights = (topk_weights / topk_weights.sum(-1, keepdim=True)).to(torch.float32)
|
||||
topk_ids = torch.stack(
|
||||
[torch.randperm(E, device=device)[:topk] for _ in range(T)]
|
||||
).to(torch.int32)
|
||||
topk_output = StandardTopKOutput(
|
||||
topk_weights=topk_weights, topk_ids=topk_ids, router_logits=None
|
||||
)
|
||||
# num_experts == num_local_experts => filter_expert=False (pure TP layout)
|
||||
cfg = MoeRunnerConfig(
|
||||
num_experts=E,
|
||||
num_local_experts=E,
|
||||
top_k=topk,
|
||||
inplace=False,
|
||||
activation="silu",
|
||||
is_gated=True,
|
||||
)
|
||||
|
||||
kwargs = dict(
|
||||
w1=w1,
|
||||
w2=w2,
|
||||
topk_output=topk_output,
|
||||
moe_runner_config=cfg,
|
||||
use_fp8_w8a8=True,
|
||||
w1_scale=w1s,
|
||||
w2_scale=w2s,
|
||||
block_shape=[128, 128],
|
||||
)
|
||||
ref = fused_experts(hidden_states=x, **kwargs)
|
||||
|
||||
q_pad, s_pad = sglang_per_token_group_quant_fp8_row_padded(x, GROUP)
|
||||
out = fused_experts(hidden_states=x, a1_q=q_pad, a1_scale=s_pad, **kwargs)
|
||||
|
||||
bitwise = torch.equal(out, ref)
|
||||
close = torch.allclose(out.float(), ref.float(), atol=1e-2, rtol=1e-2)
|
||||
_report(
|
||||
f"routed-consumer T={T} K={K} E={E} topk={topk}",
|
||||
close,
|
||||
f"(bitwise={bitwise}, max|d|={(out.float() - ref.float()).abs().max().item():.3e})",
|
||||
)
|
||||
return bitwise
|
||||
|
||||
|
||||
def main():
|
||||
assert torch.cuda.is_available(), "CUDA required"
|
||||
set_global_server_args_for_scheduler(ServerArgs(model_path="dummy"))
|
||||
device = "cuda"
|
||||
torch.manual_seed(0)
|
||||
|
||||
print("== (1) quantize-once vs routed-baseline quant equivalence ==")
|
||||
all_bitwise_q = True
|
||||
for T in (1, 3, 4093, 4096):
|
||||
for K in (6144, 7168):
|
||||
all_bitwise_q &= test_quant_equivalence(T, K, device)
|
||||
|
||||
print("== (2) shared consumer (cutlass w8a8 linear) ==")
|
||||
# N=512 mirrors a tp8 shared expert gate_up (2*2048/8); must be %128==0.
|
||||
for T in (4093, 4096):
|
||||
for K in (6144, 7168):
|
||||
test_shared_consumer(T, K, 512, device)
|
||||
|
||||
print(
|
||||
"== (2b) shared consumer (deepgemm w8a8 linear, JIT DG recommended JIT-DeepGEMM config) =="
|
||||
)
|
||||
from sglang.srt.layers import deep_gemm_wrapper
|
||||
|
||||
if not deep_gemm_wrapper.ENABLE_JIT_DEEPGEMM:
|
||||
print("SKIP: deep_gemm unavailable or SGLANG_ENABLE_JIT_DEEPGEMM=0")
|
||||
elif deep_gemm_wrapper.DEEPGEMM_SCALE_UE8M0:
|
||||
print("SKIP: Blackwell UE8M0 scale layout (gated ineligible by design)")
|
||||
else:
|
||||
for T in (4093, 4096):
|
||||
for K in (6144, 7168):
|
||||
test_shared_consumer_deepgemm(T, K, 512, device)
|
||||
# DG accepts N % 64 (cutlass needs % 128) -- exercise the DG-only shape.
|
||||
test_shared_consumer_deepgemm(4096, 7168, 320, device)
|
||||
|
||||
print("== (3) routed consumer (triton fused_experts) ==")
|
||||
# Identical in both cutlass and JIT-DG configs: the MoE runner stays
|
||||
# triton with a2a=none (is_deepgemm_moe_runner_backend_enabled() is False
|
||||
# for auto + a2a=none even when SGLANG_ENABLE_JIT_DEEPGEMM=1).
|
||||
for T in (61, 4093, 4096):
|
||||
test_routed_consumer(T, 7168, E=32, I=256, topk=8, device=device)
|
||||
|
||||
if not all_bitwise_q:
|
||||
print(
|
||||
"NOTE: quantize-once q/s not bitwise vs the routed baseline quant "
|
||||
"(AOT v2 vs JIT v2 kernel rounding) -- routed consumer is then "
|
||||
"allclose-only; document this in the PR."
|
||||
)
|
||||
if FAILURES:
|
||||
print(f"FAILED: {FAILURES}")
|
||||
sys.exit(1)
|
||||
print("ALL PASS")
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -0,0 +1,120 @@
|
||||
"""Tests for the SM90 Q8KV8 born-fp8 q-prep JIT kernel.
|
||||
|
||||
Gates (mirroring benchmark/kernels/deepseek/benchmark_q8kv8_q_prep.py conventions):
|
||||
(a) vs the Triton absorbed_bmm_concat_cast_q_fp8 "two_dot" variant with
|
||||
atol/rtol=2e-2 on the fp32 view (the accumulation order matches, so the
|
||||
output is empirically bitwise identical on SM90, but only the tolerance
|
||||
is contractual);
|
||||
(b) vs an fp64 bmm reference: mean |err| must match two_dot's;
|
||||
rope half must be bit-exact (identical bf16 -> fp8 conversion chain).
|
||||
"""
|
||||
|
||||
import pytest
|
||||
import torch
|
||||
|
||||
from sglang.test.ci.ci_register import register_cuda_ci
|
||||
|
||||
register_cuda_ci(est_time=240, stage="base-b-kernel-unit", runner_config="1-gpu-large")
|
||||
|
||||
N_LORA = 512 # kv_lora_rank
|
||||
ROPE = 64 # qk_rope_head_dim
|
||||
|
||||
|
||||
def _is_sm90() -> bool:
|
||||
return torch.cuda.is_available() and torch.cuda.get_device_capability() == (9, 0)
|
||||
|
||||
|
||||
requires_sm90 = pytest.mark.skipif(not _is_sm90(), reason="requires SM90 (Hopper)")
|
||||
|
||||
|
||||
def _make_inputs(T: int, H: int, K: int, seed: int = 1234, magnitude: float = 1.0):
|
||||
# Production layout: q_nope/q_rope are strided views of one [T, H, K+R]
|
||||
# q_b_proj output; w_kc is the N-major absorbed weight [H, K, N] with
|
||||
# strides (K*N, 1, K).
|
||||
g = torch.Generator(device="cuda").manual_seed(seed)
|
||||
q = (
|
||||
torch.randn((T, H, K + ROPE), generator=g, device="cuda", dtype=torch.float32)
|
||||
* magnitude
|
||||
).to(torch.bfloat16)
|
||||
w = (
|
||||
torch.randn((H, N_LORA, K), generator=g, device="cuda", dtype=torch.float32)
|
||||
/ K**0.5
|
||||
).to(torch.bfloat16)
|
||||
return q[..., :K], q[..., K:], w.transpose(1, 2)
|
||||
|
||||
|
||||
@requires_sm90
|
||||
@pytest.mark.parametrize("T", [1, 437, 1024])
|
||||
@pytest.mark.parametrize(
|
||||
"h_k", [(64, 192), (128, 128)], ids=["glm_h64_k192", "ds_h128_k128"]
|
||||
)
|
||||
@pytest.mark.parametrize("pad_heads_extra", [0, 2])
|
||||
def test_qprep_vs_triton_two_dot(T, h_k, pad_heads_extra):
|
||||
from sglang.kernels.ops.attention.qprep_bf16_fp8_sm90 import q8kv8_qprep_fwd
|
||||
from sglang.kernels.ops.kvcache.cache_ops import (
|
||||
absorbed_bmm_concat_cast_q_fp8,
|
||||
)
|
||||
|
||||
H, K = h_k
|
||||
q_nope, q_rope, w_kc = _make_inputs(T, H, K)
|
||||
ph = H + pad_heads_extra
|
||||
ref = torch.zeros((T, ph, N_LORA + ROPE), dtype=torch.float8_e4m3fn, device="cuda")
|
||||
out = torch.zeros_like(ref)
|
||||
absorbed_bmm_concat_cast_q_fp8(ref, q_nope, w_kc, q_rope, H, variant="two_dot")
|
||||
q8kv8_qprep_fwd(out, q_nope, w_kc, q_rope, H)
|
||||
torch.cuda.synchronize()
|
||||
|
||||
# rope half: identical conversion chain -> bit-exact.
|
||||
assert torch.equal(
|
||||
ref[:, :H, N_LORA:].contiguous().view(torch.uint8),
|
||||
out[:, :H, N_LORA:].contiguous().view(torch.uint8),
|
||||
), "rope half must be bit-exact vs the Triton kernel"
|
||||
|
||||
# gate (a): nope half within tolerance on the fp32 view.
|
||||
torch.testing.assert_close(
|
||||
out[:, :H].to(torch.float32),
|
||||
ref[:, :H].to(torch.float32),
|
||||
atol=2e-2,
|
||||
rtol=2e-2,
|
||||
)
|
||||
|
||||
# padded head slice must stay untouched.
|
||||
if pad_heads_extra:
|
||||
assert out[:, H:].view(torch.uint8).max().item() == 0
|
||||
|
||||
|
||||
@requires_sm90
|
||||
@pytest.mark.parametrize(
|
||||
"h_k", [(64, 192), (128, 128)], ids=["glm_h64_k192", "ds_h128_k128"]
|
||||
)
|
||||
def test_qprep_fp64_reference_parity(h_k):
|
||||
from sglang.kernels.ops.attention.qprep_bf16_fp8_sm90 import q8kv8_qprep_fwd
|
||||
from sglang.kernels.ops.kvcache.cache_ops import (
|
||||
absorbed_bmm_concat_cast_q_fp8,
|
||||
)
|
||||
|
||||
H, K = h_k
|
||||
T = 437
|
||||
q_nope, q_rope, w_kc = _make_inputs(T, H, K, seed=5678)
|
||||
ref8 = torch.zeros((T, H, N_LORA + ROPE), dtype=torch.float8_e4m3fn, device="cuda")
|
||||
out8 = torch.zeros_like(ref8)
|
||||
absorbed_bmm_concat_cast_q_fp8(ref8, q_nope, w_kc, q_rope, H, variant="two_dot")
|
||||
q8kv8_qprep_fwd(out8, q_nope, w_kc, q_rope, H)
|
||||
torch.cuda.synchronize()
|
||||
|
||||
# gate (b): the CUDA kernel's fp8 must land as close to the exact bmm as
|
||||
# the Triton kernel's (same quantization noise floor, ~1.8e-2 mean).
|
||||
ref64 = torch.bmm(
|
||||
q_nope.transpose(0, 1).to(torch.float64), w_kc.to(torch.float64)
|
||||
).transpose(0, 1)
|
||||
err_tri = (ref8[..., :N_LORA].to(torch.float64) - ref64).abs().mean().item()
|
||||
err_cuda = (out8[..., :N_LORA].to(torch.float64) - ref64).abs().mean().item()
|
||||
assert (
|
||||
err_cuda <= 1.05 * err_tri
|
||||
), f"CUDA fp64-ref mean |err| {err_cuda:.4e} exceeds Triton's {err_tri:.4e}"
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
import sys
|
||||
|
||||
sys.exit(pytest.main([__file__, "-v", "-s"]))
|
||||
@@ -192,6 +192,85 @@ def test_sparse_mla_q8kv8_prefill_corner_cases(
|
||||
_run_and_check(d_qk, with_sink, s_q=s_q, topk=topk, s_kv=s_kv)
|
||||
|
||||
|
||||
# topk_length WITHOUT attn_sink (the production early-exit path for
|
||||
# SGLANG_ENABLE_DSA_Q8KV8_TOPK_LENGTH): rows with a trailing -1 pad run must
|
||||
# be BITWISE identical to the full-topk dispatch that masks those pads, and
|
||||
# must match the fp32 reference on the truncated index range.
|
||||
@pytest.mark.skipif(
|
||||
not _sm90_available(), reason="Q8KV8 sparse prefill requires SM90 CUDA"
|
||||
)
|
||||
@pytest.mark.parametrize(
|
||||
"d_qk,s_q,topk,s_kv",
|
||||
[
|
||||
(576, 8, TOPK, S_KV),
|
||||
(576, 65, 256, 592),
|
||||
(512, 8, TOPK, S_KV),
|
||||
],
|
||||
)
|
||||
def test_sparse_mla_q8kv8_prefill_topk_length_only(
|
||||
d_qk: int, s_q: int, topk: int, s_kv: int
|
||||
):
|
||||
from sglang.kernels.ops.attention.sparse_mla_q8kv8_prefill_sm90 import (
|
||||
sparse_mla_q8kv8_prefill_fwd,
|
||||
)
|
||||
|
||||
q, kv, indices, sm_scale, q_scale, kv_scale, _, _ = _make_case(
|
||||
d_qk, False, s_q=s_q, topk=topk, s_kv=s_kv
|
||||
)
|
||||
# Trailing pad runs of varying size, including a 1-valid-entry row (the
|
||||
# production clamp(min=1) floor) and full rows.
|
||||
lengths = [
|
||||
topk if i % 3 == 0 else (1 if i % 3 == 1 else max(topk - 32, topk // 2))
|
||||
for i in range(s_q)
|
||||
]
|
||||
topk_length = torch.tensor(lengths, dtype=torch.int32, device="cuda")
|
||||
for q_idx, valid_topk in enumerate(lengths):
|
||||
if valid_topk < topk:
|
||||
indices[q_idx, 0, valid_topk:] = -1
|
||||
|
||||
out, max_logits, lse = sparse_mla_q8kv8_prefill_fwd(
|
||||
q=q,
|
||||
kv=kv,
|
||||
indices=indices,
|
||||
sm_scale=sm_scale,
|
||||
q_scale=q_scale,
|
||||
kv_scale=kv_scale,
|
||||
d_v=D_V,
|
||||
attn_sink=None,
|
||||
topk_length=topk_length,
|
||||
)
|
||||
out_full, max_logits_full, lse_full = sparse_mla_q8kv8_prefill_fwd(
|
||||
q=q,
|
||||
kv=kv,
|
||||
indices=indices,
|
||||
sm_scale=sm_scale,
|
||||
q_scale=q_scale,
|
||||
kv_scale=kv_scale,
|
||||
d_v=D_V,
|
||||
attn_sink=None,
|
||||
topk_length=None,
|
||||
)
|
||||
torch.cuda.synchronize()
|
||||
|
||||
assert torch.equal(out, out_full)
|
||||
assert torch.equal(max_logits, max_logits_full)
|
||||
assert torch.equal(lse, lse_full)
|
||||
|
||||
ref, ref_max_logits, ref_lse = _torch_sparse_attention_ref(
|
||||
q=q,
|
||||
kv=kv,
|
||||
indices=indices,
|
||||
sm_scale=sm_scale,
|
||||
q_scale=q_scale,
|
||||
kv_scale=kv_scale,
|
||||
attn_sink=None,
|
||||
topk_length=topk_length,
|
||||
)
|
||||
torch.testing.assert_close(out.float(), ref, atol=8e-2, rtol=8e-2)
|
||||
torch.testing.assert_close(max_logits.float(), ref_max_logits, atol=1e-2, rtol=1e-2)
|
||||
torch.testing.assert_close(lse.float(), ref_lse, atol=2e-3, rtol=2e-3)
|
||||
|
||||
|
||||
# Precision / accuracy: no-sink only because these metrics are intended to
|
||||
# approximate the current DeepSeek NSA E2E path. Sink behavior is still covered
|
||||
# above as kernel feature coverage, but sink-enabled precision numbers should
|
||||
@@ -602,5 +681,44 @@ def test_sparse_mla_q8kv8_prefill_large_skv():
|
||||
assert cos > 0.99, f"cos {cos:.4f} <= 0.99"
|
||||
|
||||
|
||||
# Backend-side topk_length derivation (backscan Triton kernel): must equal the
|
||||
# reference "last non-negative position + 1 (min 1)" on every pad pattern the
|
||||
# production topk output can produce (trailing runs), plus adversarial ones
|
||||
# (interleaved -1s, all-pad, full rows) where the trailing-run semantics still
|
||||
# define the correct consumed range.
|
||||
@pytest.mark.skipif(
|
||||
not _sm90_available(), reason="Q8KV8 sparse prefill requires SM90 CUDA"
|
||||
)
|
||||
@pytest.mark.parametrize("s_q,topk", [(437, 2048), (7, 128), (65, 256), (4096, 2048)])
|
||||
def test_q8kv8_topk_length_backscan(s_q: int, topk: int):
|
||||
from sglang.kernels.ops.kvcache.cache_ops import (
|
||||
q8kv8_topk_length_from_indices,
|
||||
)
|
||||
|
||||
generator = torch.Generator(device="cuda")
|
||||
generator.manual_seed(4000 + s_q + topk)
|
||||
indices = torch.randint(
|
||||
0, 1 << 20, (s_q, topk), dtype=torch.int32, device="cuda", generator=generator
|
||||
)
|
||||
# Row patterns: full, trailing pad runs of every length, all-pad,
|
||||
# interleaved -1s inside the valid range.
|
||||
for i in range(s_q):
|
||||
mode = i % 5
|
||||
if mode == 1:
|
||||
indices[i, max(1, i % topk) :] = -1
|
||||
elif mode == 2:
|
||||
indices[i, :] = -1
|
||||
elif mode == 3:
|
||||
indices[i, i % topk :: 7] = -1 # interleaved + trailing mix
|
||||
elif mode == 4:
|
||||
indices[i, topk - 1 :] = -1
|
||||
|
||||
got = q8kv8_topk_length_from_indices(indices)
|
||||
|
||||
ramp = torch.arange(1, topk + 1, dtype=torch.int32, device="cuda")
|
||||
ref = ((indices >= 0).int() * ramp).amax(dim=-1).clamp_(min=1)
|
||||
assert torch.equal(got, ref)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
sys.exit(pytest.main([__file__, "-v", "-s"]))
|
||||
|
||||
Reference in New Issue
Block a user