Skip to content

[AMD] [GLM5] Add opt-in Triton fp8 sparse-MLA prefill kernel for gfx950 - #28975

Merged
HaiShaw merged 1 commit into
sgl-project:mainfrom
Raiden-Makoto:RM/dsa-triton-prefill
Jun 24, 2026
Merged

HaiShaw merged 1 commit into
sgl-project:mainfrom
Raiden-Makoto:RM/dsa-triton-prefill

Conversation

@Raiden-Makoto

@Raiden-Makoto Raiden-Makoto commented Jun 23, 2026 •

Copy link
Copy Markdown
Contributor

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_fwd runs the same TileLang
partial+combine kernel for prefill as for decode. The sparse-MLA attention tile
is tiny — M = 16 heads (one 16x16 MFMA) — 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=32 tile fits
the problem shape and saturates the GPU on block count instead. It also reads
q_nope/q_rope directly, skipping the per-layer concat (the kernel splits q
into 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 is
left on TileLang.

Modifications

  • New python/sglang/srt/layers/attention/dsa/triton_sparse_mla.py: a per-query
    Triton flash kernel over the indexer-selected topk KV. Autotuned over
    BLOCK_N/num_warps/num_stages; guards an all-masked query row against
    NaN (finite softmax shift when the row has no valid key). Reads q_nope
    (width d_v) and q_rope (width dim-d_v) as two separate tensors.
  • python/sglang/srt/layers/attention/dsa_backend.py: in forward_extend, route
    the fp8 prefill path to the Triton kernel and pass q_nope/q_rope directly,
    skipping concat_mla_absorb_q_general.
    • Opt-in, default off: enable with SGLANG_DSA_TRITON_PREFILL=1.
    • Gated to gfx950 + the validated shape (num_heads==16, d_v==512,
      tail==64, topk==2048); other archs/shapes use TileLang.
    • Prefill only; decode stays on TileLang.

Accuracy Tests

GSM8K 5-shot, full 1319 questions, GLM-5.1-MXFP4, MI350X (gfx950), tp4,
fp8_e4m3 KV:

accuracy
TileLang prefill (default) 0.938
Triton prefill (SGLANG_DSA_TRITON_PREFILL=1) 0.941

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):

concurrency total tok/s TTFT (ms) ITL (ms) E2EL (ms)
2 944 920 17.74 19491
4 1717 1752 18.56 21436
8 2847 2887 20.66 25895
16 4343 5116 24.23 33946
32 6233 9588 28.68 47438
64 8158 18558 35.54 72308

This PR (SGLANG_DSA_TRITON_PREFILL=1):

concurrency total tok/s TTFT (ms) TTFT delta ITL (ms) E2EL (ms)
2 957 798 -13.3% 17.71 19228
4 1738 1511 -13.8% 18.58 21126
8 2922 2486 -13.9% 20.73 25217
16 4515 4410 -13.8% 24.22 32655
32 6574 8276 -13.7% 28.56 44814
64 8750 16051 -13.5% 35.36 67431

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

  1. Ping Merge Oncalls to start the process. See the PR Merge Process.
  2. Get approvals from CODEOWNERS and other reviewers.
  3. Trigger CI tests with comments or contact authorized users to do so.
    • Common commands include /tag-and-rerun-ci, /tag-run-ci-label, /rerun-failed-ci
  4. After green CI and required approvals, ask Merge Oncalls or people with Write permission to merge the PR.

CI States

Latest PR Test (Base): ❌ Run #27994768455
Latest PR Test (Extra): ⏳ Run #27994768350

@gemini-code-assist

Copy link
Copy Markdown
Contributor

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).
@gemini-code-assist

Copy link
Copy Markdown
Contributor

Warning

You have reached your daily quota limit. Please wait up to 24 hours and I will start processing your requests again!

@Raiden-Makoto

Raiden-Makoto commented Jun 23, 2026 •

Copy link
Copy Markdown
Contributor Author

/tag-and-rerun-ci

@HaiShaw HaiShaw left a comment

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

LGTM

@HaiShaw

HaiShaw commented Jun 24, 2026

Copy link
Copy Markdown
Collaborator

@Raiden-Makoto please add a follow-up PR to update GLM cookbook.
cc @1am9trash

@HaiShaw
HaiShaw merged commit 7454735 into sgl-project:main Jun 24, 2026
193 of 221 checks passed
Raiden-Makoto added a commit to Raiden-Makoto/squidward that referenced this pull request Jun 24, 2026
- 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.
@amd-oshkarav

amd-oshkarav commented Jun 25, 2026 •

Copy link
Copy Markdown
Contributor

Guys, @Raiden-Makoto,
I have tried the feature on GLM 5.2 FP8 which theoretically should work, but the server cannot be started when kv is fp8_e4m3. Could you check up if I do something wrong, or indeed we have an issue?
Here are the details.
Image lmsysorg/sglang-rocm:v0.5.13.post1-rocm720-mi35x-20260624 as it is,
Server start script:

#!/usr/bin/env bash
export SGLANG_USE_AITER="1"
export CUDA_VISIBLE_DEVICES="0,1,2,3,4,5,6,7"
export SAFETENSORS_FAST_GPU="1"
export SGLANG_DSA_TRITON_PREFILL="1"
python -m sglang.launch_server \
  --model zai-org/GLM-5.2-FP8 \
  --host 0.0.0.0 \
  --port 5060 \
  --tp-size 8 \
  --served-model-name "glm-5.2-fp8" \
  --trust-remote-code \
  --watchdog-timeout 12000 \
  --attention-backend dsa \
  --dsa-prefill-backend aiter \
  --dsa-decode-backend tilelang \
  --reasoning-parser glm45 \
  --tool-call-parser glm47 \
  --disable-radix-cache \
  --kv-cache-dtype fp8_e4m3 \
    2>&1 | tee serve.log
Error at the end of server start:
[AITER] /sgl-workspace/aiter/aiter_meta/csrc/py_itfs_cu/asm_mla.cu:227 mla_decode_stage1_asm_fwd: fp8 Q requires q_scale and kv_scale
Fatal Python error: Aborted

Also, when I remove --kv-cache-dtype fp8_e4m3, the server can be started and does not garbage on a simple request.

@Raiden-Makoto

Copy link
Copy Markdown
Contributor Author

@amd-oshkarav two things:

1. Wrong prefill backend. You passed --dsa-prefill-backend aiter, but the Triton prefill kernel from this PR only runs under the tilelang backend (--dsa-prefill-backend tilelang, which is the default). With aiter, the Triton kernel never engages — so SGLANG_DSA_TRITON_PREFILL=1 is a no-op in your run and this PR's code path isn't exercised at all. Use --dsa-prefill-backend tilelang (and keep --kv-cache-dtype fp8_e4m3).

2. There is a real bug, but it's separate from this PR. With --dsa-prefill-backend aiter, prefill/extend goes to _forward_aiter_extend, which calls the aiter asm MLA kernel. That code only ever sets kv_scale and leaves q_scale = None:

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 fp8 Q requires q_scale and kv_scale guard and aborts. When you drop --kv-cache-dtype fp8_e4m3 the if is skipped → no crash, which matches what you saw. This comes from the aiter DSA path added in #26639, not from this PR.

@Raiden-Makoto
Raiden-Makoto deleted the RM/dsa-triton-prefill branch July 7, 2026 19:06
Raiden-Makoto added a commit to Raiden-Makoto/squidward that referenced this pull request Jul 8, 2026
…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.
Raiden-Makoto added a commit to Raiden-Makoto/squidward that referenced this pull request Jul 8, 2026
… 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.
Raiden-Makoto added a commit to Raiden-Makoto/squidward that referenced this pull request Jul 8, 2026
…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.
Raiden-Makoto added a commit to Raiden-Makoto/squidward that referenced this pull request Jul 8, 2026
…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.
Raiden-Makoto added a commit to Raiden-Makoto/squidward that referenced this pull request Jul 8, 2026
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.
Raiden-Makoto added a commit to Raiden-Makoto/squidward that referenced this pull request Jul 8, 2026
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).
Chronostasys pushed a commit to MindLab-Research/sglang that referenced this pull request Aug 24, 2026
…50 (sgl-project#28975)

Co-authored-by: Raiden-Makoto <Raiden-Makoto@users.noreply.github.com>
fanxingran added a commit to xiaobochen-amd/sglang that referenced this pull request Sep 3, 2026
…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.
Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Labels

None yet

Projects

None yet

Development

Successfully merging this pull request may close these issues.

3 participants