Repository navigation
[Kernel] Add HIP BF16 sparse MLA for GLM-5.3-Flash H64 prefill on gfx950 - #6037
sumin-hong wants to merge 1 commit into
Conversation
Expose Full140 v1 and hybrid v2 through one CSR attention API for BF16 D512 latent K/V and H16/H64. Reuse HipKittens, support caller-owned output and split workspace, and preserve current-stream and graph semantics. Include FP64 and graph tests, dispatch-boundary coverage, a pinned-VGPR assembly check, and a reproducible same-checkout Gluon benchmark. Signed-off-by: Sumin Hong <sumin.hong@moreh.io> Co-authored-by: OpenAI Codex <codex@openai.com>
🏷️ CI GuideRuns automatically on every PR:
Extended tests (opt-in via labels):
One backend per PR: PR title tags & labels: |
Summary
Add an opt-in HIP BF16 sparse-MLA operator for GLM-5.3-Flash, targeting
H64 prefill with large query chunks on gfx950. In the causal-prefill benchmark
at Q2048/warm, v2 is 110.89–160.46% faster than tuned Gluon with
disjoint KV. Paired shared-KV measurements are 50.80–96.56% faster.
These are isolated operator measurements on MI355X.
H64 prefill is the performance target of this PR. H16 and decode/small-Q
measurements, including their regressions, are provided as supplementary
data below. The existing Gluon dispatch is unchanged; selecting this API
is opt-in.
Origin: AITER PR #3459
This work starts from ROCm/aiter PR #3459,
"Introduce the 1st Gen 64 and 128 Heads MLA Decode Kernel for DeepSeek V4 for
MI35x", at
e961cb0c0cc3d442fdabb97617d9bb438a4740da, specifically itspersistent HipKittens m16x4 implementation. We reuse its pinned-register/fragment
helpers and BF16 LDS readers. Building on that work, this PR adds the
GLM-5.3-Flash BF16 D512 shared-K/V contract with no appended RoPE, per-query
physical CSR selection, and the Full140 v1 / hybrid v2 implementations below.
H64 prefill performance
The benchmark models causal kpool4 selection with tail entries, covering
both the first chunk (prefix0) and a chunk after a 32K-token prefix. Disjoint
KV is the primary comparison; shared KV is a paired supplementary control
with identical Q/KV values and logical selections. Selection is synthetic
and uniform, rather than generated by the model's learned indexer.
At H64/Q2048/warm, v2 uses S1 for every row below. Each implementation's
split is selected in the first run and held fixed for an independent
reverse-order repeat. All comparison columns use Faster (%):
100 × (baseline time / v2 time − 1). Positive values indicate a gain;negative values indicate a regression. This is an operator speed ratio,
not an end-to-end serving-throughput measurement.
For these H64 rows, Gluon uses S2 for the first chunk and S1 with a 32K
prefix. The H64 causal-prefill measurements cover seven Q values, paired KV
layouts, two prefixes, warm/cold cache and both output contracts:
4,160 timing records / 249,600 event samples, with zero accuracy exclusions.
They were measured on the submission head/build10_rebase. The supplementary
H16 measurements bring the total to 8,416 records / 504,960 samples.
The timing includes attention main and split reduction. Indexer, KV writes,
CSR conversion and the rest of the model are outside this measurement;
end-to-end prefill latency or throughput gains have not been established.
H64 prefill bandwidth — Q2048, warm
These are the same independently repeated points as the headline table. Cells
show median µs / logical GB/s (split); the final column is Faster (%)
for v2 versus tuned Gluon. V1 is included as the Full140 reference.
Logical bytes count BF16 Q and O, one gathered BF16 KV vector per valid
selection, the full fixed-width CSR indices/indptr and optional FP32 LSE.
For E valid selections, D512 and top-k width2051:
B = 4QHD + 2DE + 4Q×2051 + 4(Q+1) [+ 4QH for LSE];logical GB/s = B / median_us / 1000.Repeated head reads, split partials and other intermediate traffic are excluded.
The same B is used for every variant at a given input/output contract.
This is useful logical bandwidth, not a hardware-counter measurement of HBM traffic.
Shared KV can reuse cache lines across queries, and warm runs can reuse data
across replays; neither is assigned an HBM-only lower bound.
H64 prefill bandwidth — Q2048, disjoint KV, cold
Cold uses a 512 MiB flush outside the timing event. Useful QK/PV work is
4EHD; padding, softmax, address arithmetic and split reduction FLOPs are excluded.With a 32K prefix, cold/disjoint v2 reaches 3,972–4,035 logical GB/s,
versus 1,560–1,572 GB/s
for tuned Gluon. Because the logical-byte numerator is identical, the bandwidth
ratio expresses the measured latency gain; it does not independently establish
a reduction in physical traffic or identify the source of stalls.
H64 prefill across chunk sizes — disjoint KV, warm
All result columns report Faster (%) versus tuned Gluon.
Large chunks provide the strongest gains. Small first-chunk shapes regress;
the all-Q first-chunk warm GM Faster (%) is +9.64 to +9.95
with disjoint KV and -0.75 to -0.51 with shared KV. The measured gain depends on the chunk
and KV geometry; these results do not define an automatic routing policy.
Operator and implementations
aiter.sparse_mla_bf16_fwdconsumes BF16 queries and a BF16 latent usedfor both K and V, with D512, no appended RoPE, and H16/H64: the absorbed attention
geometry used by GLM-5.3-Flash. It accepts per-query CSR global slot ids,
masks invalid slots and preserves the full selected set, including a
2051-entry tail.
One
aiter.sparse_mla_bf16_fwdAPI exposes two reproducible implementations:The caller can reuse contiguous O, optional natural-log LSE and normalized
FP32 split workspace. The torch-free ctypes entry uses the current stream;
mutation/fake registration makes the operation visible to
torch.compile.Metadata checks do not copy CSR values to the host during graph capture.
The module reuses AITER's pinned HipKittens dependency and CK-free headers,
with FP32 denormal preservation and an explicit pinned-VGPR build check.
The existing Gluon dispatch remains unchanged. The new API defaults to v2
and explicit S=1; these defaults are not a claim of the optimal split for
every workload. Engine selection and fallback belong to a subsequent vLLM
integration, using this API without copying the kernel implementation.
Validation
Submission base
e0cb2ddbde2e7e71c1fe2c8d05f3f6f656360e14, implementationea21fa0a6c5d40916d2a7112812727f54217ebd7, MI355X/gfx950.The fixed-K appendix tables were measured at base
80a3b0b09448bb5b8ade9a5946bec7549601331a, implementation1121378ad62b1d1e1babe366c3fcb5ed2d6d14cb(build09).Rebasing across seven upstream commits changed none of the 14 PR files or
the Gluon/JIT/header dependencies used here. A fresh JIT build (build10_rebase)
has identical preprocessed device source and assembly after normalizing only
the generated HIP CUID symbol. The full 514-test suite and 336 native trace
dispatch checks pass again on that build. The original fixed-K sweep was not repeated
after the rebase; the causal-prefill results above were measured on
build10_rebase. Trace durations are excluded from performance tables.
HipKittens
d3cd9b31cb0ff611ff64b5701f57ccdeb7712f39is the existing AITER pin.PyTorch
2.12.0+git6bbd260, Triton3.7.1+gitf0b55c07, HIP7.2 runtime,ROCm7.2.3 compiler; one otherwise idle GPU.
Black 26.3.0 and clang-format 18.1.8 pass on the changed files. Ruff 0.15.7
passes on all four new Python modules; the touched
aiter/__init__.pyhasthe same 46 pre-existing F403 diagnostics as the submission base, with no
new diagnostics. The PR contains one signed-off commit with the HIP operator,
API, tests and benchmark. The Gluon address-control patch is separate diagnostic
evidence and is not included in this PR.
Reproduce the operator tests and the supplemental fixed-K sweep:
The sweep uses disjoint 32K-slot pools per query, H16/H64,
Q=1/32/128/256/512/1024/2048, K=2048/2051, and S=1/2/4/8/16/32.
Every point has 60 graph/event samples and 10 warmups. The 512 MiB cold flush
is outside the timing event. Native versions share input/output/workspace
addresses; Gluon shares inputs/final O and keeps its own partial layout.
Trace durations are not used as timing measurements.
Model accuracy and integration
Supplemental eager GLM-5.3-Flash TP4/H16 checks exercised both versions on
short, 32K and 128K prompts. All 11 sparse layers reached the native API in
prefill and decode, with no unsupported fallback. FP64 checks passed on
385 sampled query rows for v1 and 374 for v2; the existing attention path
also passed on those same Q/KV inputs. Maximum native normalized RMS error
was 0.00281, and maximum LSE absolute error was 5.73e-6. These are sampled
operator checks inside the model. Baseline greedy A/A varied before native
execution; token identity is not an acceptance requirement.
GSM8K — GLM-5.3-Flash
Local OFF/ON scores use lm-eval 5-shot strict-match on all 1,319 questions,
4×MI355X TP4/H16, BF16 KV, batch4, temperature1.0/top_p0.95, a 32,768-token
output limit and max reasoning effort, with the same weights, prompts and
generation settings (seed0). ON uses this PR's original v2/S1 kernel.
The published reference evaluated 1,319 questions on 4×GB300 with SGLang;
its hardware, backend and evaluator differ, so it is a reference point.
Scores vary between runs, including with the unchanged kernel. The table
shows the existing full OFF result and the latest full v2 rerun; it does not
establish a statistically significant accuracy improvement or equivalence.
The latest ON evaluation resumed from saved batches after model restarts,
counting each completed batch once. GPQA-Diamond, RULER and v1 full-model
scores remain pending.
The new API defaults to v2/S1; tuned results require the displayed explicit
split. A future vLLM PR should call this operator without copying kernel
source, retain fallback, validate physical CSR mapping and graph workspace
reuse, and measure the actual serving path. Broader model accuracy and
end-to-end serving speedups, especially for TP4/H16, remain unestablished.
All validation is associated with the source/binary hashes above. A later
rebase or dispatch change requires the relevant checks again. Full wheel
installation and Inductor validation remain outside the completed scope.
Related work
Reviewed at the following pinned revisions on 2026-09-30; PR statuses below
refer to that review.
fragment/register helpers and BF16 LDS readers. That V4 FP8+RoPE decode
interface differs from this BF16 D512 per-query CSR interface.
#5551, merged: the same-checkout
Gluon sparse-MLA implementation is the primary direct comparison. It has
broader format/geometry coverage. Both its shipped policy and explicit
split sweep are included in the measured comparison.
847d0148:gfx942 support is complementary to this gfx950-only HIP implementation.
c69f539b:related BF16 D512 sparse-prefill work for gfx942, including attention sinks
and low-head-count tuning. It targets a different backend/device path;
its reported performance is not a denominator for this gfx950 operator.
9b5cf23e:FP8 KV, separate RoPE, H1–16; its reported numbers are not a same-contract
BF16 D512 comparison and are not used as a speedup denominator here.
Supplementary performance data
H16 causal-prefill reference, including all Q2048 warm regressions
H16 is supported as a reference path. V1 uses one wave / 33,024 B LDS and
the original reducer; v2 uses four waves / 68,352 B LDS and the unrolled
reducer. The causal sweep adds 4,256 timing records / 255,360 event samples
for H16, bringing the H64+H16 total to 8,416 records / 504,960 samples.
At Q2048/warm, v2 Faster (%) versus tuned Gluon ranges from
-39.78 to -11.21 across these prefill conditions. GLM-5.3-Flash TP4 uses H16; the H64 prefill headline
therefore does not demonstrate a serving speedup for the TP4 model.
H16/Q2048/warm reference, including regressions:
Gluon uses S2 for the first chunk, except H16/shared/O-only uses S4;
all long-prefix entries use S1. V2 uses S1 throughout this table.
Decode / small-Q reference and fixed-K aggregate results
The fixed-K operator sweep includes the small-Q shapes relevant to decode.
It does not measure end-to-end decode latency or batched serving throughput.
Its all-Q geometric means are supplementary to the causal-prefill headline
and include every measured Q, including regressions.
Splits below are selected in the first run and fixed for the independent
reverse-order repeat. Comparison columns report Faster (%), computed as
100 × (GM(baseline time / v2 time) − 1), with regressions shown as negativevalues. Each summary gives 14 Q×K shapes equal weight; it is not a
serving-workload average or an end-to-end throughput measurement.
H64 fixed-K reference — Faster (%)
For H64/Q2048/K2048/warm, O-only is 165.22% faster than tuned
Gluon: 1108.085 µs (v2) versus 2938.834 µs (Gluon), all S1; v1 is
1436.547 µs. O+LSE is 163.33% faster: 1094.447 µs versus
2881.978 µs; v1 is 1421.368 µs.
H64 small-Q regressions remain: O-only warm Q1 and Q32 have Faster (%)
of -13.34 and -7.70 respectively (GM over both K values).
H16 reference — Faster (%)
H16 Faster (%) is negative versus tuned Gluon across all measured warm Q values.
For O+LSE/Q2048/K2048/warm, v2 is 825.385 µs versus 733.724 µs for Gluon.
Full shape/cache/contract tables, with H64 first and H16 reference second,
follow below. These results do not support replacing Gluon across all shapes.
The trace observes one CTA per query for v2 H64/S1 and four for Gluon;
H16 uses one in both. This validates the dispatch/head-sharing distinction.
It does not establish physical HBM reread counts or their latency share.
Gluon address-boundary exclusions and corrected-baseline control
Baseline accuracy qualification
Original Gluon H64/Q2048/S16 and S32 fail the sampled FP64 gate for both K
values. Separate address-only controls identify the partial-buffer resource
boundary: the retained LLVM IR range is 2 GiB−2 bytes. An int64 global-pointer
control and a query-rebased buffer control both pass all nine probe cases;
their sampled BF16 outputs agree bitwise. At exactly 2 GiB, two final output
elements can be wrong; at 4 GiB, later queries are affected.
H64/Q1024/S32 reaches the same boundary but can pass the tolerance gate.
Its four raw timing points per phase are also excluded after this diagnosis.
The original raw records are retained. All best/first-selected results are
unchanged; matched-split coverage accounts for these additional exclusions.
Supplemental timing of the query-rebased Gluon control covers H64 Q1024/2048,
K2048/2051 and every split, both output contracts and independent repeats: 608
timing records, 36,480 samples, zero accuracy exclusions. S1 remains its best
split in every shape/cache/contract. With first-selected splits fixed, v2 vs
this corrected Gluon has the following Faster (%) (GM over four Q×K shapes):
The controls are separate lab modules, not changes to the measured upstream
baseline or a general Gluon fix included in this PR. Other formats/architectures
would need their own correctness validation before an upstream addressing fix.
Complete fixed-K independent-repeat tables: small-Q/decode reference, all other shapes, H16, cache and output contracts, bandwidth and roofline
AITER sparse MLA: independent performance results
Measured AITER base
80a3b0b09448bb5b8ade9a5946bec7549601331a, measured source1121378ad62b1d1e1babe366c3fcb5ed2d6d14cb.v1 is Full140+U at H64 and the legacy one-wave implementation at H16. v2 is the hybrid+U implementation. This measures the isolated attention main+reducer, not a model or serving engine.
Each query owns a disjoint 32K-slot pool. H16/H64, Q=1/32/128/256/512/1024/2048, K=2048/2051, S=1/2/4/8/16/32. Native input/output/workspace addresses are shared within each case; Gluon shares inputs/final output and retains its native partial layout.
The tables use the reverse-order repeat, with each variant's split fixed from its first run. Each point has 60 graph/event samples. Warm/cold are separate; a 512 MiB cold flush is outside the timing event. The shipped Gluon policy is reported separately from its tuned split. Native API S=1 is an explicit default, so these tuned results require the displayed split.
The audit checks completed coverage, source/binary/input hashes, accuracy and sampled KFD ownership. All regressions and accuracy exclusions remain in JSON/CSV. Same-split and retrospective best-in-run results are separate in the underlying summaries.
These fixed-K results are supplementary to the H64 causal-prefill headline. All comparison columns use Faster (%) = 100 × (baseline time / v2 time − 1); negative values indicate regressions. GM summaries apply this formula to the geometric mean of 14 equally weighted Q×K speed ratios, not a serving-workload average. The timing, bandwidth and split cells retain their original measured values.
O-only: 2,112 timing records / 126,720 event samples across both runs.
After the address controls, four raw tolerance-passing Q1024/S32 timings per phase are additionally excluded: 2,104 eligible timings per output contract. The raw records remain unchanged. All best/first-selected results are unchanged; matched-split coverage uses the corrected exclusions.
O+LSE: 2,112 timing records / 126,720 event samples across both runs.
After the address controls, four raw tolerance-passing Q1024/S32 timings per phase are additionally excluded: 2,104 eligible timings per output contract. The raw records remain unchanged. All best/first-selected results are unchanged; matched-split coverage uses the corrected exclusions.
H64 — fixed-K reference
O-only: Faster (%) (GM)
O-only: H64, K2048, warm
Cells: median µs / logical GB/s (split S).
O-only: H64, K2051, warm
Cells: median µs / logical GB/s (split S).
O-only: H64, K2048, cold
Cells: median µs / logical GB/s (split S).
O-only: H64, K2051, cold
Cells: median µs / logical GB/s (split S).
O+LSE: Faster (%) (GM)
O+LSE: H64, K2048, warm
Cells: median µs / logical GB/s (split S).
O+LSE: H64, K2051, warm
Cells: median µs / logical GB/s (split S).
O+LSE: H64, K2048, cold
Cells: median µs / logical GB/s (split S).
O+LSE: H64, K2051, cold
Cells: median µs / logical GB/s (split S).
H16 — reference results
O-only: Faster (%) (GM)
O-only: H16, K2048, warm
Cells: median µs / logical GB/s (split S).
O-only: H16, K2051, warm
Cells: median µs / logical GB/s (split S).
O-only: H16, K2048, cold
Cells: median µs / logical GB/s (split S).
O-only: H16, K2051, cold
Cells: median µs / logical GB/s (split S).
O+LSE: Faster (%) (GM)
O+LSE: H16, K2048, warm
Cells: median µs / logical GB/s (split S).
O+LSE: H16, K2051, warm
Cells: median µs / logical GB/s (split S).
O+LSE: H16, K2048, cold
Cells: median µs / logical GB/s (split S).
O+LSE: H16, K2051, cold
Cells: median µs / logical GB/s (split S).
Roofline reference: Q2048/K2048 cold
AMD's MI355X datasheet gives per-GPU theoretical ceilings of 8 TB/s HBM and 2.5166 PFLOP/s dense BF16. These are advertised ceilings, not measured sustained rates. This kernel does not use hardware structured sparsity.
Logical bytes count Q, one gathered KV copy, CSR indices, O and optional LSE. Repeated reads and partial buffers are excluded. Logical bandwidth is a useful-byte proxy, not a physical HBM counter or utilization measurement. QK/PV TFLOP/s excludes softmax, address arithmetic and other scalar/vector work. The ideal model is memory-limited at these intensities; the actual bottleneck and extra traffic remain unmeasured. Warm cache hits can invalidate an HBM-only lower bound, hence this table uses cold points.
Interpretation and routing boundary
The H16 and H64 results must remain separate. The measured H64 medium/large-Q gains support an opt-in HIP path; they do not support replacing Gluon for every head count or query shape. The existing Gluon dispatcher remains unchanged. A future vLLM adapter must explicitly select supported geometry and an appropriate split, retain fallback, and separately validate its physical CSR mapping and serving behavior.
At Q2048/S1, the retained trace observes one CTA per query for v2 H64 versus four for Gluon H64. At H16 both use one. This supports the source-level distinction in head sharing; it does not measure physical KV rereads or establish how much of the latency difference they cause.
The raw tolerance gates exclude H64/Q2048/S16 and S32. Address controls additionally invalidate H64/Q1024/S32 at the 2 GiB partial boundary, even when its tolerance gate passes. The original errors and controls are retained separately. Gluon's shipped and best valid paths at Q1024/Q2048 use S1, so the primary comparisons retain full shape coverage.
AI assistance
OpenAI Codex assisted with implementation, tests, analysis and documentation.
The submitter is responsible for reviewing the code and the stated evidence.