Skip to content

[FlyDSL] Single-launch Mega-mHC on gfx950 - #6179

Merged
junhaha666 merged 14 commits into
mainfrom
anguyenh/flydsl-mega-mhc
Oct 8, 2026
Merged

junhaha666 merged 14 commits into
mainfrom
anguyenh/flydsl-mega-mhc

Conversation

@anhminhnguyenhoang

@anhminhnguyenhoang anhminhnguyenhoang commented Oct 6, 2026 •

Copy link
Copy Markdown
Contributor

Motivation

FlyDSL implementation of the DeepSeek-V4.1 delayed (Single-Pass) mHC boundary for gfx950. One launch does post-mix, input mix with the carried pre, RMSNorm, gate projection, Sinkhorn and next pre, and writes the bf16 norm plus, optionally, SGLang's per-32 ue8m0 FP8 quantization of it. It replaces the two Triton launches of mhc_fused_post_pre_delayed_rmsnorm (#5824) and SGLang's seam including its separate norm + fake-quant kernel.

Wiring into SGLang's hc_boundary_fused is out of scope; all speedups are isolated kernel times.

Technical Details

Files Changed

  • aiter/ops/flydsl/__init__.py: lazy export of flydsl_mega_mhc.
  • aiter/ops/flydsl/kernels/mega_mhc.py: kernel, check_config, ue8m0_quant.
  • aiter/ops/flydsl/mega_mhc_kernels.py: wrapper (Triton-seam signature + out_dtype, config=), policy get_mega_mhc_config, per-stream scratch, fn pre-pack cache.
  • op_tests/test_flydsl_mega_mhc.py: correctness, perf vs the Triton seam, config, streams, DIST_FINISH, LATE_DESC, graph capture and ue8m0 edge tests.

Kernel Architecture

  • Policy picks BLOCK_M, warps, NUM_KSPLIT, TILE_K from T, out_dtype and CU count. Default: 16 tokens per workgroup, 8 column-split warps; fn pre-packed as bf16 hi/lo in MFMA B order.
  • Order: post-mix → collapse → sum of squares + gate projection → split-K reduce → rstd → Sinkhorn → layer input.
  • Split-K (up to 20 splits, all on one XCD): the last-arriving workgroup finishes. Agent-scope counter atomic with workgroup-scope release / acquire fences.
  • DIST_FINISH: the last arrival publishes rstd, each split rescales its own x1 from LDS. Residency guard plus bounded spin and hand-off: never hangs.
  • Prefill: X1_LDS_SLOTS (x1 in LDS), wave-quantization split-K, SEG128 128 B-line streams with nt stores and loads. Decode: DPP Sinkhorn shuffles, LATE_DESC.
  • out_dtype: bf16 → norm; fp8_grid → (norm, bf16 grid); mxfp8 → (norm, fp8 codes, e8m0). Each finish path quantizes the unit it stores (DPP-quad amax, v_cvt_scalef32_pk_fp8_bf16). bf16 ISA is unchanged by the ue8m0 code.
  • CUDA graphs: scratch and fn pack are filled by eager calls only; a capture that would need them raises RuntimeError.

Limitations

  • gfx950 only; thresholds tuned at H = 5120, 256 CUs. T · 4 · H · 2 < 2^31.
  • DIST_FINISH assumes co-resident splits; under heavy concurrent CU use it falls back to hand-off (correct, ≥ 0.5 ms per stuck round). Off with config={"DIST_FINISH": False} or a CU mask.
  • Output choice follows SGLang v0.5.21 + PR #42055 per side and consumer backend: attention mxfp8 for MXFP8 wqkv_a, else fp8_grid; FFN mxfp8 with SGLANG_HIP_FFN_NORM_MXFP8, else bf16.
  • Groups holding inf / NaN are unspecified (SGLang clamps first); other groups are unaffected.
  • ue8m0 cost over bf16: +0.9 µs at T = 384, +12.6% at T = 16384 for fp8_grid (+7.3% with the consumer GEMM).
  • No 128-wide-group FP8 (128×128-block checkpoints).

Test Plan

  • python op_tests/test_flydsl_mega_mhc.py (gfx950; picked up by aiter_test.sh).
  • Torch reference for post, no_post, identity_pre, all out_dtypes, T = 0..16400; ue8m0 outputs bit-exact against the torch rule on the kernel's own norm, every finish path.
  • Crafted quantizer groups (amax = 448·2^k, zeros, subnormals, near-max, inf / NaN).
  • Against SGLang v0.5.21 on identical inputs, and its in-model seam at TP4 (external scripts, outputs quoted below).
  • Coherence and stress probes, graph capture, two streams, CU-masked run; ISA / VGPR / scratch; black, ruff.

Test Result

Correctness:

  • Op test: 0 failures (all modes, dtypes, knob sets, capture, streams, deadlock).
  • vs SGLang v0.5.21 (M = 4 … 16384, post, no_post, first block): norm and residual bitwise or within one bf16 ulp (≥ 99.996% bitwise); gates ≤ 5.8e-4; ue8m0 outputs mismatch only where the norm differs by one ulp (≤ 0.009% of groups). [Triton/Gluon] [gfx950] [dsv4.1-flash] mHC fused kernel #5824 gives the same errors.
  • Edge test: 0 of 1120 crafted groups differ.
  • bf16 ISA byte-identical across the ue8m0 change; no spills or scratch.

Performance vs SGLang v0.5.21 (kernel device time, both measured the same way):

M Form Attention, fp8_grid FFN, bf16 mxfp8
4 post 13.6 → 7.6 µs (1.79×) 13.8 → 7.4 µs (1.86×) 13.6 → 7.7 µs (1.78×)
4 no_post 12.5 → 6.4 µs (1.95×) 12.5 → 6.4 µs (1.97×) 12.7 → 6.5 µs (1.95×)
64 post 15.9 → 9.0 µs (1.76×) 15.8 → 8.6 µs (1.83×) 15.8 → 8.9 µs (1.77×)
1024 post 63.7 → 29.5 µs (2.15×) 59.7 → 28.5 µs (2.10×) 63.2 → 30.0 µs (2.11×)
16384 post 767.4 → 411.3 µs (1.87×) 728.0 → 378.0 µs (1.93×) 754.8 → 395.9 µs (1.91×)
16384 no_post 670.4 → 256.7 µs (2.61×) 624.6 → 228.2 µs (2.74×) 652.7 → 239.5 µs (2.73×)

In-model estimate (SGLang v0.5.21 + PR #42055, TP4, 16,384-token chunks, decode bs = 4: 57.6 ms per chunk, 0.99 ms per step): FlyDSL isolated, same forms and outputs, gives 31.1 ms per chunk (−26.5 ms) and 0.51 ms per decode step (−0.47 ms). These combine in-model SGLang times with isolated FlyDSL times.

Performance comparison, bf16 post, H = 5120 (runtime / bandwidth / efficiency vs 5.86 TB/s; FlyDSL and Triton in CUDA graph, SGLang device time):

Implementation T=1 T=32 T=384 T=2048 T=4096 T=4160 T=8192 T=16384
SGLang seam, 3 kernels 15.7 us / 0.007 / 0.1% 19.4 us / 0.17 / 2.9% 37.2 us / 1.06 / 18.0% 118.5 us / 1.77 / 30.2% 232.0 us / 1.81 / 30.9% 236.1 us / 1.80 / 30.8% 429.4 us / 1.95 / 33.3% 797.7 us / 2.10 / 35.9%
aiter Triton seam, 2 kernels 10.8 us / 0.009 / 0.2% 11.9 us / 0.28 / 4.7% 21.3 us / 1.85 / 31.5% 61.6 us / 3.40 / 58.1% 126.4 us / 3.32 / 56.6% 141.5 us / 3.01 / 51.4% 244.8 us / 3.43 / 58.5% 431.8 us / 3.89 / 66.3%
FlyDSL Mega-mHC, 1 kernel 7.5 us / 0.014 / 0.2% 8.3 us / 0.39 / 6.7% 13.8 us / 2.85 / 48.6% 43.7 us / 4.80 / 81.9% 102.5 us / 4.09 / 69.8% 124.4 us / 3.42 / 58.4% 199.6 us / 4.20 / 71.7% 372.8 us / 4.50 / 76.8%
  • 2–2.7× faster than SGLang's seam, 1.1–1.5× faster than the Triton seam; 70–77% of the copy ceiling at large T.
  • T = 4160 dips to 58%: 260 blocks on 256 CUs leave a near-empty second round; the split-K fallback recovers 179 → 124 µs.
  • Small T is latency-bound.

Realistic Bandwidth Ceiling

At 16k rows a post seam moves 1.678 GB (1.845 GB with the fp8 grid). A plain copy of those bytes reaches 5.86 TB/s, so ≈ 286 µs is the floor for one kernel; FlyDSL is at 1.32× of it, or 1.13× of a 5 TB/s byte roofline (copy at T = 4096, 419 MB: 71.6 µs).

@github-actions

github-actions Bot commented Oct 6, 2026

Copy link
Copy Markdown
Contributor

🏷️ CI Guide

Runs automatically on every PR:

  • ✅ Pre-checks (submodule verification, code formatting)
  • ✅ Aiter op tests (gfx942 + gfx950)
  • ✅ Triton tests on MI35X (only when aiter/ops/triton/** or related paths are changed)

Extended tests (opt-in via labels):

Label Tests
ci:gfx1250-ffm-triton Run the five-shard gfx1250 FFM Triton test suite
ci:triton-300x Run an additional Triton test job on MI300X in PRs (added automatically when gfx942 configs change); main branch always runs both MI35X and MI300X
ci:triton-355 Run the full Triton test suite on MI35X, not only the tests the change affects
multigpu Aiter multi-GPU tests on the 8-GPU runner
ci:sglang SGLang integration tests: DeepSeek-R1-MXFP4 accuracy, Qwen 3.5 accuracy
ci:atom ATOM benchmark: DeepSeek-R1-0528, GPT-OSS-120B
ci:atom_full ATOM accuracy suite for PR and main models from ATOM models_accuracy.json
ci:vllm vLLM benchmark: GPT-OSS-120B, DeepSeek-R1-0528, Kimi-K2.5
ci:all All standard extended tests (excludes ci:atom_full)

Only add ci:atom_full for FlyDSL or Triton upgrades.
Add labels via the sidebar or gh pr edit 6179 --add-label <label>

One backend per PR:
A PR changes one kernel backend: [Triton/Gluon] (Triton and Gluon count as one), [HIP], [ASM], [CK], [OPUS] or [FlyDSL]. If the title ends up with two backend tags, split the PR -- as stacked pull requests when one part cannot merge without the other.

PR title tags & labels:
Component tags ([Triton/Gluon], [HIP], [CK], [ASM], ...) are added to the PR title and as PR labels automatically from the changed files and re-synced on every push — change-type tags like [fix]/[Perf], op tags like [MLA], and human labels (ci:*) are left untouched. Add the no-auto-title label to stop the title rewrites; labels stay in sync either way.

@anhminhnguyenhoang anhminhnguyenhoang changed the title flydsl mega mhc [FlyDSL] Single-launch Mega-mHC on gfx950 Oct 6, 2026
anhminhnguyenhoang and others added 14 commits October 7, 2026 13:55
Adds flydsl_mega_mhc, a single-launch FlyDSL kernel for the DeepSeek-V4.1
delayed (Single-Pass) mHC seam: post-mix, new residual, collapse with the
carried pre gate, gate projection (bf16 hi/lo MFMA on a pre-packed fn),
RMS, Sinkhorn and the next block's RMSNorm input in one launch.

- Same API as mhc_fused_post_pre_delayed_rmsnorm, plus out_dtype
  ("bf16" or "fp8": e4m3, group-32 fp32 scales) and a config override.
- Split-K finish by the last workgroup of each token block; splits are
  mapped to one XCD (COHERENCE="xcd"), with an agent-scope mode kept.
- Modes: post, no_post (Engram seam), identity_pre.
- HIP-style (arch, cu_num) launch policy from a gfx950 sweep.
- op_tests/test_flydsl_mega_mhc.py: candidates vs the fp32 reference and
  the Triton seam, knob sweep with ISA resources, two-stream concurrency,
  CUDA-graph replay.

MI355X, H=5120, post mode: T=4096 124 us bf16 / 118 us fp8 vs 135 us for
the Triton seam; T=1 8.7 us vs 11.8 us.

Co-authored-by: Cursor <cursoragent@cursor.com>
CUDA graphs (correctness): a call during capture could record the fn
pack and the counter memset instead of running them, and cache the
unwritten results, so a later eager call or another graph read garbage;
growing the split-K scratch freed buffers that captured graphs still
launch on. Now a call during capture that would pack fn or allocate
scratch raises RuntimeError (warm up on the capture stream first), no
cache is written during capture, and superseded scratch is retained.
A warmed replay stays a single kernel.

xcd mode: workgroup-scope release/acquire fences around the split-K
counter (the agent-scope atomic is kept), so the ordering no longer
relies on s_waitcnt + barrier alone. Final ISA is unchanged.

op_test: test_mega_mhc_capture (cold capture, eager call after capture,
two graphs with scratch growth replayed in reverse); the graph check
warms on the capture stream.

Co-authored-by: Cursor <cursoragent@cursor.com>
…it (P1 sweep)

Co-authored-by: Cursor <cursoragent@cursor.com>
…tays off)

Co-authored-by: Cursor <cursoragent@cursor.com>
… finish (P3), fill-aware token-kernel rule

Co-authored-by: Cursor <cursoragent@cursor.com>
…walk (P4)

Co-authored-by: Cursor <cursoragent@cursor.com>
…e finding)

Co-authored-by: Cursor <cursoragent@cursor.com>
… guard and hand-off (D2)

Co-authored-by: Cursor <cursoragent@cursor.com>
…rnarg reorder (D3)

Co-authored-by: Cursor <cursoragent@cursor.com>
…split and policy (G1)

Co-authored-by: Cursor <cursoragent@cursor.com>
Remove knobs that are fixed or never selected by the launch policy:
COHERENCE (agent mode deleted; XCD mapping whenever NUM_KSPLIT > 1),
FN_PREPACKED (in-kernel fp32 fn split deleted), SINKHORN_RCP,
NT_STREAMS (superseded by NT_LD / NT_ST), WARPS_PER_SIMD (constant
waves_per_eu = 2), PERSIST_PREFETCH (always on with PERSIST_WGS) and
SHUFFLE_DPP = 2 (SHUFFLE_DPP is now a bool).

Also delete the dead load_fn_all helper, merge the identical FP8 q store
cache bits into cm_st, raise ValueError on unknown config= keys, trim the
policy docstring to its rules and drop references to out-of-tree notes.

All 36 policy snapshot cases (T = 1..16400, bf16/FP8, post / no_post /
identity_pre, forced hand-off) give bit-identical outputs and identical
final ISA; CUDA-graph timing matches the previous numbers within noise.

Co-authored-by: Cursor <cursoragent@cursor.com>
Kernel (final ISA identical on all 36 policy snapshot cases, outputs
bit-identical):
- MFMA via fx.make_mma_atom + fx.mma_atom_call instead of the raw
  rocdl.mfma_f32_16x16x32_bf16 intrinsic.
- Stream loads via fx.rocdl.raw_ptr_buffer_load, stores via
  buffer_ops.buffer_store(offset_is_bytes=True); drops the direct
  flydsl._mlir imports and the private _RAW_PTR_BUFFER_AUX_IS_ATTRIBUTE read.
- fx.max instead of fx.maximumf, shared LOG2E from kernels_common,
  typed loop-carried state without raw ir_value() wraps.

Test: drop the unmeasured 6.5 TB/s "floor us" column; T = 0 rows record
only the error (nothing to time).

Co-authored-by: Cursor <cursoragent@cursor.com>
… amax/448 FP8 mode (U1)

out_dtype bf16 | fp8_grid -> (norm, grid) | mxfp8 -> (norm, codes, e8m0): the per-32
ue8m0 rule of SGLang's fake_quant_fp8_activation on the bf16 norm, written by every finish
path through one store_unit (DPP-quad amax, v_cvt_scalef32 fp8 <-> bf16). bf16 kernels
compile to the same ISA; the amax/448 mode and its FP8-only policy branches are removed.
Tests check the ue8m0 outputs bit-exact on every finish path and on crafted groups.

Co-authored-by: Cursor <cursoragent@cursor.com>
…ge test, NT_LD per out_dtype)

- remove PERSIST_WGS / FN_EARLY (no policy path sets them), the scaled_cvt fallback and
  the probe hook; fx.gemm instead of mma_atom_call (bf16 ISA byte-identical)
- input validation raises ValueError / TypeError instead of assert; get_mega_mhc_config
  validates out_dtype; ue8m0 dtypes check for v_cvt_scalef32; explicit float8_e4m3fn
- op test: test_mega_mhc_ue8m0_edges (crafted + inf / NaN groups), Triton seam baseline
  column for every out_dtype, one default config-sweep row
- NT_LD threshold per out_dtype: 230 MB for fp8_grid (paired win at T = 2048 / 2304,
  isolated and with the consumer GEMM)

Co-authored-by: Cursor <cursoragent@cursor.com>
@anhminhnguyenhoang
anhminhnguyenhoang marked this pull request as ready for review October 7, 2026 18:58
@anhminhnguyenhoang
anhminhnguyenhoang requested review from coderfeli and removed request for vgokhale October 8, 2026 07:29

@coderfeli coderfeli 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.

@anhminhnguyenhoang

Copy link
Copy Markdown
Contributor Author

mHC seam: best HIP vs Triton vs SGLang vs FlyDSL Mega-mHC (MI355X, bf16, post, n=4, H=5120)

Time per seam in µs (lower is better). HIP, Triton and FlyDSL: CUDA-graph replay, one process, interleaved, median of 5. SGLang: kernel device time (its HIP prefill kernel can't be graph-captured).

T SGLang (3 kernels) Best HIP (aiter mhc_fused_post_pre) Triton (aiter #5824, 2 kernels) FlyDSL Mega-mHC (1 kernel) FlyDSL vs SGLang vs HIP vs Triton
1 15.7 7.7 (forced) 10.5 7.3 2.15× 1.05× 1.44×
32 19.4 8.6 (forced) 11.4 8.2 2.37× 1.05× 1.39×
384 37.2 22.7 (forced) 21.0 13.8 2.70× 1.64× 1.52×
2048 118.5 69.6 (forced) 60.7 40.7 2.91× 1.71× 1.49×
4096 232.0 127.5 (forced / large_m) 126.6 102.2 2.27× 1.25× 1.24×
8192 429.4 270.7 (large_m) 245.2 200.1 2.15× 1.35× 1.23×
16384 797.7 542.2 (auto) 434.6 380.3 2.10× 1.43× 1.14×

What each one computes:

  • SGLang (39bcc19): the shifted seam, in three kernels: boundary, reduce + Sinkhorn, then RMSNorm.
  • Best HIP: the fastest of the forced, auto and large_m variants at each T, all with packed BF16 weights. It uses the unshifted ordering and has no RMSNorm, so it does less work than the other three and isn't usable as-is for DSv4.1.
  • Triton: mhc_fused_post_pre_delayed_rmsnorm, the shifted seam with RMSNorm.
  • FlyDSL: flydsl_mega_mhc at 9288b1996f, the shifted seam with RMSNorm.

Notes:

  • SGLang is the device time measured on Oct 2 (FlyDSL at 7ff033ba28 then). The HIP, Triton and FlyDSL columns were re-measured together on Oct 8.
  • T ≤ 32 is latency-bound.

@anhminhnguyenhoang
anhminhnguyenhoang requested a review from a team October 8, 2026 10:49
@junhaha666
junhaha666 merged commit 31aef27 into main Oct 8, 2026
96 checks passed
@junhaha666
junhaha666 deleted the anguyenh/flydsl-mega-mhc branch October 8, 2026 13:29
Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Labels

Projects

None yet

Development

Successfully merging this pull request may close these issues.

3 participants