Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
2 changes: 1 addition & 1 deletion README.md
Original file line number Diff line number Diff line change
Expand Up @@ -138,7 +138,7 @@ env `why`, gap notes): flags, `snake_case` names, file names, `org/repo` paths a
- [Kimi-K3 on MI355X](kimi_k3_playbook.md) — Day-0 hybrid KDA/MLA MoE, MXFP4; plain + DSpark speculative decoding, GSM8K/AIME25 ([`test_kimi_k3.sh`](test_kimi_k3.sh), [`eval_kimi_k3.sh`](eval_kimi_k3.sh)) — plus a tuned recipe worth **+27%** throughput and **+68%** on DSpark, found by the search harness in [`grid_k3/`](grid_k3/README.md). Re-measured 2026-08-08 on the upstream `rocm/sgl-dev:v0.5.16` image (§2.0), which is ~10% faster and needs no fork or patch, and **corrected**: the long-context DSpark cliff was a draft-checkpoint RoPE bug (§5.4a), not a limit of speculative decoding. Extra harnesses from that pass: [`aime26_eval.py`](aime26_eval.py) (AIME26 is not in `run_eval`; inherits SGLang's AIME25 scorer), [`degeneracy_probe.py`](degeneracy_probe.py) (why `--dataset-name random` overstates accept length), [`prep_speedbench.py`](prep_speedbench.py) and [`prep_agentic_trace.py`](prep_agentic_trace.py)
- [GLM-5.2-FP8 on MI300X](glm52_fp8_playbook.md) — DSA tilelang, FP8; GSM8K/AIME25 + long-context ([`test_glm52_fp8.sh`](test_glm52_fp8.sh))
- [GLM-5.2-FP8 on MI355X](glm52_fp8_mi355x_playbook.md) — the gfx950 re-measurement on SGLang 0.5.17 / ROCm 7.2.4, which retires the two mandatory `bpreshuffle` patches (the CK fix the old cell named as its own exit condition has landed), adds **MTP/NEXTN speculative decoding** — it works on ROCm now — and **`fp8_e4m3` KV**, which is legal on the tilelang DSA path on ROCm and roughly doubles the pool. Fills the `balanced` and `high-throughput` cells; rows generated by [`gen_glm52_mi355x_rows.py`](gen_glm52_mi355x_rows.py)
- [GLM-5.3-Flash on MI355X](glm53_flash_playbook.md) — 320B/18B hybrid KDA+DSA model on the stacked SGLang model/ROCm PRs, with FP8 KV, AITER MoE/mHC, fused k-pool, GSM8K/AIME25, and an apples-to-apples GLM-5.2 comparison. Rows generated by [`gen_glm53_mi355x_rows.py`](gen_glm53_mi355x_rows.py). §11 is the pitfalls list — including a node where a broken `rocminfo` silently made `tilelang` target `gfx900` and `flydsl` target `gfx942`, which brings the server up healthy and then kills it on the first real request ([`rocminfo_shim.sh`](glm53_flash/rocminfo_shim.sh))
- [GLM-5.3-Flash on MI355X](glm53_flash_playbook.md) — 320B/18B hybrid KDA+DSA model on the stacked SGLang model/ROCm PRs, with FP8 KV, AITER MoE/mHC, fused k-pool, GSM8K/AIME25, and an apples-to-apples GLM-5.2 comparison. Rows generated by [`gen_glm53_mi355x_rows.py`](gen_glm53_mi355x_rows.py). §11 is the pitfalls list — including a node where a broken `rocminfo` silently made `tilelang` target `gfx900` and `flydsl` target `gfx942`, which brings the server up healthy and then kills it on the first real request ([`rocminfo_shim.sh`](glm53_flash/rocminfo_shim.sh)). §12 is why `aiter#5069`'s measured **-25%** kernel win moved serving throughput by **0.0%**: the tables it retuned are read by two tiny GEMMs, half of it is unreachable on a block-quantised checkpoint, and our own pinned table covers only M=1 and M=32 while prefill drives M to 8192 ([`coverage_report.py`](glm53_flash/coverage_report.py) is the census tool)
- [GLM-5.3 on MI355X](glm53_fp8_mi355x_playbook.md) — the FULL GLM-5.3, not Flash: `glm_moe_dsa` at 756 GB, with the same architecture, quantization layout, shard count, and safetensors byte count as GLM-5.2-FP8 at the pinned revisions, so its recipes transfer. It now has a visible `not_benchmarked` site cell: serving and the long-cold-prefill fix were verified, while the single-run door-side throughput remains excluded from the datasheet. On the affected AITER `c16d44b9` + Triton 3.7 stack, a single long prompt not already in the prefix cache aborts every TP rank on a Triton/LLVM `iota_range(Begin <= End)` assertion. The cause is AITER's `fp8_mqa_logits` switching, above 2 GiB of logits, to a `gl.store` path that does not compile for its `BLOCK_M = 2` shape — so the single-request wall is `min(L, chunked_prefill_size) * L * 4 >= 2**31` when `num_q > 4096`, putting it at 32,767 tokens for chunk size 16,384 and at 23,170 when a long prompt fits in one chunk. Older AITER and Triton 3.6 direct probes do not reproduce it, so this is version-specific rather than a property of every gfx950 DSA deployment. It is easy to miss in agentic traffic because a long agentic turn is usually a long cache *hit*. §1 has the one-unit boundary measurements, version matrix, wall table, and upstream fix, with which 1,001,869 cold tokens answer
- [DeepSeek-V4-Flash-0731 on MI355X](dsv4_flash_playbook.md) — official 284B/13B checkpoint, FP4 experts + block-FP8 dense weights, target-only `unified_kv_triton`, three complete GSM8K runs and median-of-3 fixed-shape serving results ([`test_dsv4_flash.sh`](test_dsv4_flash.sh), rows generated by [`gen_dsv4_mi355x_rows.py`](gen_dsv4_mi355x_rows.py))
- [DeepSeek-V4-Pro-0813 on MI355X](dsv4_pro_playbook.md) — official 1.6T/49B checkpoint on one 8x MI355X node, including the 20-minute startup characteristic, three complete GSM8K runs and median-of-3 serving results ([`test_dsv4_pro.sh`](test_dsv4_pro.sh), shared launcher [`test_dsv4.sh`](test_dsv4.sh))
Expand Down
76 changes: 76 additions & 0 deletions glm53_flash/coverage_report.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,76 @@
#!/usr/bin/env python3
"""Which tuned GEMM tables does this model actually consult, and how far does
padded_M drift from the real M?

Section 12 of glm53_flash_playbook.md is the write-up; this is the tool. It
answers the question an op-level speedup claim cannot: does the model read the
table that changed at all, and at which shapes.

Usage:

# launch the server with the lookup logger on
AITER_LOG_TUNED_CONFIG=1 python3 -m sglang.launch_server ... > server.log
# drive a representative load, then
python3 coverage_report.py server.log

AITER's lookup is lru_cached, so it logs once per distinct shape key: this is a
coverage census, not a call census. For time-weighting, join the reported
kernelName against a torch-profiler capture.

Caveat: AITER_LOG_TUNED_CONFIG only instruments the a16w16 (BF16) path. The
block-scale FP8 and fused-MoE tables have their own tables and their own logs;
absence of a table here is not proof it went unconsulted -- check which merged
CSVs the process materialised under /tmp/aiter_configs/ as well.
"""
import re, sys, collections
from pathlib import Path

HIT = re.compile(
r"shape is M:(?P<M>\d+), N:(?P<N>\d+), K:(?P<K>\d+).*?"
r"found padded_M: (?P<pM>\d+).*?is tuned on cu_num = (?P<cu>\d+) in "
r"(?P<file>\S+?), libtype is (?P<lib>\w+)")
MISS = re.compile(
r"shape is M:(?P<M>\d+), N:(?P<N>\d+), K:(?P<K>\d+).*?"
r"not found tuned config in (?P<file>\S+?),.*?using (?P<lib>\w+) solution")

hits, misses = [], []
for line in Path(sys.argv[1]).read_text(errors="replace").splitlines():
m = HIT.search(line)
if m:
hits.append(m.groupdict()); continue
m = MISS.search(line)
if m:
misses.append(m.groupdict())

def base(f): return f.rsplit("/", 1)[-1]

print(f"{len(hits)} distinct shapes hit, {len(misses)} distinct shapes missed\n")
Comment on lines +45 to +47

Copy link
Copy Markdown

Choose a reason for hiding this comment

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

issue (bug_risk): The report counts every matching log line as a distinct shape, but AITER's cache is per worker process, so a multi-rank server emits the same shape once per rank. The displayed hit/miss totals and percentages are therefore inflated by worker duplication rather than representing distinct shapes.

Triggers: When the server log combines lookup output from multiple tensor-parallel worker processes.

Suggested fix: Deduplicate records by the lookup key, such as (M, N, K, table, padded_M, library), before computing counts and percentages.

Suggested change
def base(f): return f.rsplit("/", 1)[-1]
print(f"{len(hits)} distinct shapes hit, {len(misses)} distinct shapes missed\n")
def base(f): return f.rsplit("/", 1)[-1]
def key(h): return (h["M"], h["N"], h["K"], h["file"], h.get("pM"), h["lib"])
hits = list({key(h): h for h in hits}.values())
misses = list({key(h): h for h in misses}.values())
print(f"{len(hits)} distinct shapes hit, {len(misses)} distinct shapes missed\n")

print("=== distinct shapes per table ===")
tbl = collections.Counter(base(h["file"]) for h in hits)
tbl_m = collections.Counter(base(h["file"]) for h in misses)
for f in sorted(set(tbl) | set(tbl_m)):
print(f" {f:<48} hit {tbl.get(f,0):>4} miss {tbl_m.get(f,0):>5}")

print("\n=== on a hit, how far padded_M is inflated over the real M ===")
buckets = collections.Counter()
worst = []
for h in hits:
M, pM = int(h["M"]), int(h["pM"])
r = pM / M
worst.append((r, M, pM, int(h["N"]), int(h["K"]), h["lib"]))
buckets["1.00 (exact)" if r == 1 else
"<=1.25" if r <= 1.25 else
"<=1.5" if r <= 1.5 else
"<=2.0" if r <= 2.0 else ">2.0"] += 1
for k in ("1.00 (exact)", "<=1.25", "<=1.5", "<=2.0", ">2.0"):
if buckets.get(k):
print(f" {k:<13} {buckets[k]:>4} shapes ({buckets[k]/len(hits):>5.1%})")
worst.sort(reverse=True)
print("\n worst 8:")
for r, M, pM, N, K, lib in worst[:8]:
print(f" M={M:<6} -> padded {pM:<6} ({r:.2f}x) N={N:<6} K={K:<6} {lib}")

print("\n=== missed shapes by (N,K) ===")
for (n, k, f), c in collections.Counter(
(h["N"], h["K"], base(h["file"])) for h in misses).most_common(10):
print(f" N={n:<6} K={k:<6} {c:>5} distinct M {f}")
16 changes: 14 additions & 2 deletions glm53_flash/setup_pr.sh
Original file line number Diff line number Diff line change
Expand Up @@ -27,9 +27,21 @@ set -euo pipefail
cd /sgl-workspace/sglang
echo '--- before ---'
git log -1 --format='%H %ci %s'
git fetch --no-tags origin 'pull/${SGLANG_PR}/head'
git merge-base --is-ancestor '${SGLANG_HEAD}' FETCH_HEAD
# The assertion here used to be 'the measured commit is still an ancestor of the
# PR head'. That broke on 2026-08-31 when #36507 was rebased: the commit is fine,
# the branch just no longer descends from it, and setup failed at step one.
# Fetching the exact object pins the tree just as tightly and survives a rebase.
git fetch --no-tags origin '${SGLANG_HEAD}' 2>/dev/null \
|| git fetch --no-tags origin 'pull/${SGLANG_PR}/head'
git cat-file -e '${SGLANG_HEAD}^{commit}' || {
echo 'FATAL: commit ${SGLANG_HEAD} is unreachable. The PR branch was rebased and the'
echo ' old commit garbage-collected. Re-measure against a current head rather'
echo ' than substituting one -- the published numbers are tied to this tree.'
exit 1
}
git checkout -q --detach '${SGLANG_HEAD}'
git merge-base --is-ancestor '${SGLANG_HEAD}' FETCH_HEAD 2>/dev/null \
|| echo 'note: measured commit is no longer an ancestor of the PR head (rebased upstream); the tree checked out above is still exactly the measured one'
Comment on lines +34 to +44

Copy link
Copy Markdown

Choose a reason for hiding this comment

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

nitpick (bug_risk): When fetching SGLANG_HEAD succeeds, FETCH_HEAD points to that same commit, so git merge-base --is-ancestor '${SGLANG_HEAD}' FETCH_HEAD compares a commit with itself and always succeeds. A rebased PR is therefore never reported by the supposedly informational ancestry check on the normal successful path.

Triggers: When the exact measured commit can be fetched directly from the remote, including the rebased-PR case described in the change.

Suggested fix: Fetch or save the PR head separately and compare SGLANG_HEAD against that commit instead of comparing it with the exact-object fetch's FETCH_HEAD.

Suggested change
git fetch --no-tags origin '${SGLANG_HEAD}' 2>/dev/null \
|| git fetch --no-tags origin 'pull/${SGLANG_PR}/head'
git cat-file -e '${SGLANG_HEAD}^{commit}' || exit 1
git checkout -q --detach '${SGLANG_HEAD}'
git merge-base --is-ancestor '${SGLANG_HEAD}' FETCH_HEAD 2>/dev/null \
|| echo 'note: measured commit is no longer an ancestor of the PR head (rebased upstream); the tree checked out above is still exactly the measured one'
git fetch --no-tags origin 'pull/${SGLANG_PR}/head' 2>/dev/null
PR_HEAD="\$(git rev-parse FETCH_HEAD)"
git fetch --no-tags origin '${SGLANG_HEAD}' 2>/dev/null \
|| git fetch --no-tags origin 'pull/${SGLANG_PR}/head'
git cat-file -e '${SGLANG_HEAD}^{commit}' || exit 1
git checkout -q --detach '${SGLANG_HEAD}'
git merge-base --is-ancestor '${SGLANG_HEAD}' "\$PR_HEAD" 2>/dev/null \
|| echo 'note: measured commit is no longer an ancestor of the PR head (rebased upstream); the tree checked out above is still exactly the measured one'

echo '--- after ---'
git log -1 --format='%H %ci %s'

Expand Down
144 changes: 144 additions & 0 deletions glm53_flash_playbook.md
Original file line number Diff line number Diff line change
Expand Up @@ -385,3 +385,147 @@ serving record's `server_info`. `bench_one_batch_server` does not emit
(batch size, output length, the three-repeat set) and rely on directory
provenance for which server produced them. Keep latency runs inside the same
tagged results directory as the serving runs, or that link is lost.

## 12. Why a 25% kernel win can be a 0% serving win

`ROCm/aiter#5069` retuned the GLM-5.2 a8w8 and BF16 GEMM configs for gfx950 and
reported, from its own measurement, 49 shapes going 3606.4us -> 2702.1us
(**-25.08%**), median +22.76% per shape, zero regressions. We A/B'd it on one
8x MI355X node with the same image, SGLang worktree, recipe and bench protocol,
changing **only the four tuning CSVs**:

| conc | GLM-5.2 delta | GLM-5.3 delta |
|---:|---:|---:|
| 1 | +0.06% | -0.01% |
| 8 | -0.05% | +0.05% |
| 16 | -0.02% | +0.08% |
| 32 | -0.05% | -0.00% |
| 64 | -0.09% | +0.15% |

Ten points, all inside a 0.04-0.35% noise floor. Both arms passed the GSM8K
gate. Before reading anything into a null result, we falsified it twice:

- **The arms really were different.** Each arm records the sha256 of what it
deployed: a8w8 `b453...` vs `a361...`, bf16 `c84f...` vs `01da...`.
- **The change really did engage.** GLM-5.2's BF16 lookup misses fell
**1256 -> 616**, exactly the `N=256, K=6144` half that the PR added rows for.

So the tuning worked and the serving throughput did not move. That is not bad
luck; it is structural.

### 12.1 Almost nothing in the model reads the tuned-GEMM table

`aiter.tuned_gemm.tgemm` has three call sites in SGLang, and one of them
(`kernels/ops/attention/dsv4/gemm.py`) is CUDA-only, i.e. dead on ROCm. For a
GLM-5.x FP8 checkpoint the live ones are:

| Module | Shape | Route |
|---|---|---|
| MoE router / gate | `N = n_routed_experts`, `K = hidden` | `tgemm.mm` |
| DSA indexer `weights_proj` (no `quant_config`, so bf16) | `N = n_heads`, `K = hidden` | `tgemm.mm` |

Everything else goes elsewhere: `qkv_a` / `q_b` / `o_proj` / `gate_up` /
`down_proj` / shared experts / KDA projections all land on
`aiter_w8a8_block_fp8_linear` -> `gemm_a8w8_blockscale*`; routed experts go to
`aiter.fused_moe` and `tuned_fmoe.csv`; `lm_head` is a plain `torch.matmul` in
`logits_processor.py`. Three different tuning tables, and the recipe's headline
GEMMs are in none of the ones this PR touched.

The serving logs agree exactly: across both A/B arms the only `(N,K)` pairs that
ever reach the BF16 table are `(256, 6144)` and `(32, 6144)` -- the router and
the indexer. Both are tiny.

Worse, half the PR is unreachable in this configuration. The non-block-scale
`gemm_a8w8_bpreshuffle` that `a8w8_bpreshuffle_tuned_gemm_glm5.2.csv` feeds is
only reached when `SGLANG_USE_AITER_FP8_PER_TOKEN` is set. GLM-5.2-FP8 is block
quantized (`weight_block_size: [128, 128]`), so it takes `gemm_a8w8_blockscale*`
instead. **Ninety-nine of the PR's 189 changed gfx950 rows are never read.**

### 12.2 The tuner optimises M values that serving never produces

Which shapes get tuned is decided entirely by the `*_untuned_gemm*.csv` shape
list. Those lists are a geometric ladder:

```
bf16 : 1 2 4 8 16 24 32 48 64 96 128 192 256 384 512 768 1024 1536 2048 3072 4096 8192 16384 32768
a8w8 : 1 2 4 8 16 32 64 128 256 512 1024 2048 4096 8192 16384 32768
```

Serving produces dense, arbitrary M -- chunked prefill and continuous batching
do not round to powers of two:

```
320 384 448 640 704 768 832 896 960 1024 1216 1984 6016 6528 6848 7104 7109 7168 7448 ...
```

Lookup does soften this. `get_GEMM_A16W16_config` probes three times: exact M,
then `getPaddedM(M,N,K,0)` (round up to a multiple of 16/32/64/128 by size
band), then `getPaddedM(M,N,K,1)` (`nextPow2`), then gives up. So `M=6016` runs
the kernel tuned for `M=8192`, and anything just past a power of two can be
served by a config tuned for nearly twice its size.

AITER already ships the fix for the shape list: `AITER_TUNE_GEMM=1` appends
every shape a real workload actually executes to the untuned CSV. A ladder of
powers of two is what you get when that step is skipped.

### 12.3 Even a perfect GEMM win is capped by what GEMM costs

A torch-profiler capture of the published GLM-5.3 cell (concurrency 32, ISL
8192) is worth reading carefully, because the raw numbers lie.
`aiter::cross_device_reduce_2stage` appears to own **97.4%** of GPU time -- but
6 of its 182 calls account for 98.9% of that, the longest running 2.35 seconds,
against a median of 282us. Those are ranks waiting at the collective, not
reducing. With the waits removed:

| Bucket | Share of real GPU compute |
|---|---:|
| TP all-reduce (genuine) | **29.9%** |
| MoE (`tuned_fmoe.csv`) | 18.8% |
| GEMM, **all** backends including blockscale | 17.3% |
| Other | 11.8% |
| mHC fusion | 10.1% |
| KDA linear attention | 6.1% |
| DSA attention | 1.7% |

Even if the retuned shapes were the entire GEMM bucket, 25% off 17.3% is 4.3%
end to end. They are instead the router and indexer slivers inside it.

### 12.4 What the coverage census actually found

`AITER_LOG_TUNED_CONFIG=1` logs the row each lookup resolves to. Running it
under a representative load on the published GLM-5.3 recipe:

| | |
|---|---:|
| Distinct shapes that hit the BF16 table | 104 |
| Distinct shapes that missed | **1616** |
| Misses that fell through to plain `torch` `F.linear` | 1608 |
| Hits at exactly the tuned M | 69.2% |
| Hits padded up, by up to 2.0x | 30.8% |

The gfx950 half of the `glm53_bf16_tuned_gemm.csv` this cookbook pins covers
**M = 1 and M = 32 only** -- precisely the two captured decode graph tiers.
Chunked prefill drives M to 8192. So the BF16 path runs unoptimised PyTorch for
the overwhelming majority of distinct shapes, and the interesting work is not
"retune the rows we have" but "the table is nearly empty for this workload".

### 12.5 What to demand of a tuning claim

An op-level speedup is a hypothesis about serving, not evidence of it. Before
believing one, ask for three things:

1. **Route proof** -- does the model read the table that changed? Check the
call site, not the filename. `/tmp/aiter_configs/` shows which merged tables
the process actually materialised; `AITER_LOG_TUNED_CONFIG=1` shows which
rows resolve. Note the merge is keyed on `(gfx, cu_num, M, N, K, ...)` and
drops the model name entirely, so a row tuned for one model will be selected
for another whenever the shape matches.
2. **Mechanism proof** -- did the change engage? A hit/miss census before and
after, so a null result cannot be confused with a no-op deployment.
3. **End-to-end A/B** -- same node, same image, same recipe, one variable, three
repeats, correctness gate before timing, and a significance bar set from the
arms' own spread. `verify.sh` and the harness conventions in section 9 are
the shape of it.

Report the null results too. A 25% kernel win that does not move serving is a
useful fact about where the time actually goes.