Repository navigation
[AMD] [GLM5] Add opt-in Triton fp8 sparse-MLA prefill kernel for gfx950 - #28975
Conversation
|
Warning You have reached your daily quota limit. Please wait up to 24 hours and I will start processing your requests again! |
The HIP fp8 sparse-MLA prefill on gfx950 runs the same TileLang partial+combine kernel as decode. Its attention tile is tiny (M=16 heads = one 16x16 MFMA), so the 256-thread TileLang block over-parallelizes it and pays intra-block coordination overhead. A per-query Triton flash kernel with a small 2-warp / BLOCK_N=32 tile fits the problem shape and saturates the GPU on block count instead. Adds triton_sparse_mla.py (autotuned over BLOCK_N/num_warps/num_stages, with an all-masked-row NaN guard) and routes the fp8 prefill path through it from forward_extend. The kernel reads q_nope/q_rope directly, skipping the per-layer concat (it splits q into main/tail internally anyway). Opt-in (default off): enable with SGLANG_DSA_TRITON_PREFILL=1. Gated to the validated shape (num_heads==16, d_v==512, tail==64, topk==2048) on gfx950; everything else falls back to TileLang. Decode is untouched. GLM-5.1-MXFP4 / MI350X (gfx950) / TP4 / fp8 KV: ~13-14% lower median TTFT across a 2-64 concurrency sweep, median ITL/E2EL flat-to-better, GSM8K 0.941 (vs 0.938 baseline; no regression).
1116c71 to
41b57e7
Compare
|
Warning You have reached your daily quota limit. Please wait up to 24 hours and I will start processing your requests again! |
|
/tag-and-rerun-ci |
|
@Raiden-Makoto please add a follow-up PR to update GLM cookbook. |
- Add the amd/GLM-5.1-MXFP4 checkpoint (tp=4, --kv-cache-dtype fp8_e4m3) as the recommended MI355X (gfx950) path in the deploy generator + ROCm command section, and document the opt-in SGLANG_DSA_TRITON_PREFILL=1 prefill kernel (follow-up to sgl-project#28975). - EAGLE speculative decoding: the old cookbook said it was unsupported on AMD and the generator gated it off for all AMD. We tested EAGLE on gfx950 (MI355X) and it works (lossless, large ITL/throughput win); gfx942 (MI300X/MI325X) is unverified. So the generator now emits the EAGLE flags by default and excludes ONLY MI300X/MI325X (gfx942); the static MXFP4 command and the AMD note reflect this. gfx942 verification can be a follow-up.
|
Guys, @Raiden-Makoto, Also, when I remove |
|
@amd-oshkarav two things: 1. Wrong prefill backend. You passed 2. There is a real bug, but it's separate from this PR. With q_scale = None
kv_scale = None
if kv_cache.dtype == fp8_dtype:
kv_scale = torch.ones((), dtype=torch.float32, device=q_kernel.device)GLM-5.2-FP8's MLA query is fp8, so the aiter kernel hits its |
…rf config) Profiling showed both are needed-on: triton sparse-MLA prefill 492ms vs tilelang 772ms, and INT4 quick-reduce all-reduce ~250ms vs 590ms nccl.
… as best target MoE/dense GEMM tuning is ceiling-limited (t64/t128 are M-buckets at equal per-token cost; dense ~37% roofline is Tensile's ceiling). allreduce + MLA-fp8 bmm are Jacob's. Real open target: our _sparse_mla_fwd (sgl-project#28975) 494ms, memory-bound with ~200ms overhead above the ~1.8ms/call BW floor.
…a target Microbench: seq=8192 kernel 1.732ms/call; logical gather BW 5.58 TB/s = 123% of measured HBM copy peak (4.55 TB/s) -> overlapping-topk KV served from L2, cache/BW-bound near ceiling, M=16 MFMA caps compute. Gluon rewrite ~10-15% at most. Retract the earlier ~200ms-headroom claim.
…her prefetch Re-measured: achievable HBM only ~3.5-4.5 TB/s (copy 4.48/read 3.52/scale 3.96), so kernel is DRAM-BW-bound at effective ~4 TB/s, not at ceiling. Compute floor 0.225ms vs 1.732ms kernel = 7.7x -> big compute slack plain Triton doesn't hide. gluon async double-buffered gather + coalesced page load is the lever to recover part of the 84ms gap.
Automates the kernel dev loop for sgl-project#28975: bf16-reference correctness plus latency/effective-bandwidth across GLM-5.2 prefill M-buckets. Run before/after a kernel change (e.g. gluon port) and diff.
Grind result: gluon kernel made correct (maxdiff 4e-4) after clearing the CDNA4 fp8 compiler wall (mfma_scaled + [32,32,64]). But 30ms (tl_dot) / 92ms (direct dot loads) vs triton 2.0ms -> 15-45x slower; gluon MFMA-layout machinery is wrong for small-M memory-bound gather. triton sgl-project#28975 near-optimal (~4.8 TB/s); 84ms vs B200 is hardware. Gluon kept flag-gated (default off).
…50 (sgl-project#28975) Co-authored-by: Raiden-Makoto <Raiden-Makoto@users.noreply.github.com>
…ackend triton Builds on the opt-in Triton sparse-MLA prefill kernel added in sgl-project#28975 and makes it usable on the serving path: a real backend choice instead of an environment variable, and an inner loop rewritten around what gfx950 actually charges for. The rewrite is extensive enough that git renders the file as a delete plus an add rather than a rename; triton_sparse_mla_prefill.py is the continuation of triton_sparse_mla.py, renamed because the decode kernel lands beside it later. Selection. `--dsa-prefill-backend triton` joins the other DSA backends in DSA_CHOICES instead of hiding behind SGLANG_DSA_TRITON_PREFILL. Construction requires gfx950 + fp8 KV. Per-request shape gates -- 8 or 16 heads, d_v=512, rope tail 64, topk 2048, 512 to 32768 tokens -- fall back to TileLang rather than failing the step. The HIP KV layout selector lists triton alongside tilelang and aiter, since it reads the same raw nope(512)+rope(64) fp8 pool. Decode is untouched, and `--dsa-decode-backend triton` is rejected outright rather than silently ignored. This kernel packs every head into one program per token, so the grid is the token count: prefill hands it thousands, decode hands it the batch size, and it measures 0.48x against TileLang at 64 heads and small batches. Use it with `--dsa-decode-backend tilelang`. Kernel. Workgroups land on XCD (pid % N_XCD), so each XCD is handed one contiguous run of tokens; the split gives the first (n_tok % N_XCD) XCDs one extra token, which keeps the mapping a bijection when the division is not exact. Softmax runs entirely in log2 space, which drops the log2(e) multiply tl.exp emits per element, turns the -inf mask into an fma addend, and folds log2(fp8_max) into the running max so the fp8 scale cancels in the final divide, removing the per-element rescale of the [H_PAD, D_V] accumulator. The denominator is reciprocated once instead of dividing the accumulator. KV loads carry no mask: page ids are clamped first, so invalid lanes read real fp8 that is then forced to -inf and multiplies out of the PV dot. Top-k indices for the next block are prefetched, which costs 3 VGPRs and buys 7% at one wavefront per SIMD. KV offsets widen to 64-bit past the int32 wrap threshold, since pool size tracks free HBM and 4.0M tokens at DIM=576 wraps the offset negative. The config is fixed at BLOCK_N=64, num_warps = clamp(H_PAD // 16, 1, 4), num_stages=1 rather than autotuned. On Triton 3.7.0 one point of the natural grid, (BLOCK_N=128, num_warps=1, num_stages=2), aborts the compiler with an LLVM assertion -- SIGABRT, not an exception autotune can catch and skip. Measured on MI355X against TileLang at the production index distribution: 1.79x on the kernel over a 100k-prompt chunk sweep, 1.95x at seq=32768. Serving A/B at 32768x28 requests, concurrency 14: TTFT p50 1.20x, TPOT p50 1.21x, throughput 1.20x. AgentX 3600 s replay at CONC=14: +9.8% P90 interactivity, TTFT p90 -14.6%, cluster throughput flat. Adds 12 tests and 6 subtests: causal-prefix padding, fully padded and single-valid-key rows, interleaved padding, the XCD bijection at divisible and non-divisible token counts, and the 64-bit KV offset path either side of the int32 threshold.
Motivation
This affects HIP serving of DSA (DeepSeek Sparse Attention) models; it was found
and validated on GLM-5.1-MXFP4 (which uses the DSA indexer) on MI350X.
On the HIP fp8 path,
tilelang_sparse_fwdruns the same TileLangpartial+combine kernel for prefill as for decode. The sparse-MLA attention tile
is tiny —
M = 16heads (one16x16MFMA) — so the 256-thread (4-warp)TileLang block over-parallelizes it and spends most of its time on intra-block
coordination rather than the matmuls (profiling at the prefill config on gfx950
shows the kernel VALU-bound at ~42% with MFMA at ~12% and ~37% occupancy).
A per-query Triton flash kernel with a small
2-warp/BLOCK_N=32tile fitsthe problem shape and saturates the GPU on block count instead. It also reads
q_nope/q_ropedirectly, skipping the per-layer concat (the kernel splits qinto main/tail internally anyway, so combining them first is wasted work — and
the HIP path uses a plain
torch.cat, not the fused CUDA concat). Decode isleft on TileLang.
Modifications
python/sglang/srt/layers/attention/dsa/triton_sparse_mla.py: a per-queryTriton flash kernel over the indexer-selected topk KV. Autotuned over
BLOCK_N/num_warps/num_stages; guards an all-masked query row againstNaN(finite softmax shift when the row has no valid key). Readsq_nope(width
d_v) andq_rope(widthdim-d_v) as two separate tensors.python/sglang/srt/layers/attention/dsa_backend.py: inforward_extend, routethe fp8 prefill path to the Triton kernel and pass
q_nope/q_ropedirectly,skipping
concat_mla_absorb_q_general.SGLANG_DSA_TRITON_PREFILL=1.num_heads==16,d_v==512,tail==64,topk==2048); other archs/shapes use TileLang.Accuracy Tests
GSM8K 5-shot, full 1319 questions, GLM-5.1-MXFP4, MI350X (gfx950), tp4,
fp8_e4m3 KV:
SGLANG_DSA_TRITON_PREFILL=1)No regression.
Speed Benchmarks
E2E
sglang.bench_serving(random, input 8192 / output 1024, median latencies),GLM-5.1-MXFP4 / MI350X / tp4 / fp8_e4m3 KV.
Baseline (TileLang prefill):
This PR (
SGLANG_DSA_TRITON_PREFILL=1):Median TTFT improves 13.3-13.9% across the sweep; median ITL is unchanged
(within +/-0.5%) and median E2EL is flat-to-better, confirming the change is
prefill-side with no decode regression.
Checklist
Review and Merge Process
/tag-and-rerun-ci,/tag-run-ci-label,/rerun-failed-ciCI States
Latest PR Test (Base): ❌ Run #27994768455
Latest PR Test (Extra): ⏳ Run #27994768350