Repository navigation
[FlyDSL] Single-launch Mega-mHC on gfx950 - #6179
Conversation
🏷️ CI GuideRuns automatically on every PR:
Extended tests (opt-in via labels):
One backend per PR: PR title tags & labels: |
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>
4ba4a94 to
9288b19
Compare
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).
What each one computes:
Notes:
|
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 nextpre, and writes the bf16 norm plus, optionally, SGLang's per-32 ue8m0 FP8 quantization of it. It replaces the two Triton launches ofmhc_fused_post_pre_delayed_rmsnorm(#5824) and SGLang's seam including its separate norm + fake-quant kernel.Wiring into SGLang's
hc_boundary_fusedis out of scope; all speedups are isolated kernel times.Technical Details
Files Changed
aiter/ops/flydsl/__init__.py: lazy export offlydsl_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=), policyget_mega_mhc_config, per-stream scratch,fnpre-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
BLOCK_M, warps,NUM_KSPLIT,TILE_KfromT,out_dtypeand CU count. Default: 16 tokens per workgroup, 8 column-split warps;fnpre-packed as bf16 hi/lo in MFMA B order.DIST_FINISH: the last arrival publishesrstd, each split rescales its ownx1from LDS. Residency guard plus bounded spin and hand-off: never hangs.X1_LDS_SLOTS(x1 in LDS), wave-quantization split-K,SEG128128 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).bf16ISA is unchanged by the ue8m0 code.fnpack are filled by eager calls only; a capture that would need them raisesRuntimeError.Limitations
H = 5120, 256 CUs.T · 4 · H · 2 < 2^31.DIST_FINISHassumes co-resident splits; under heavy concurrent CU use it falls back to hand-off (correct, ≥ 0.5 ms per stuck round). Off withconfig={"DIST_FINISH": False}or a CU mask.mxfp8for MXFP8wqkv_a, elsefp8_grid; FFNmxfp8withSGLANG_HIP_FFN_NORM_MXFP8, elsebf16.inf/NaNare unspecified (SGLang clamps first); other groups are unaffected.T = 384, +12.6% atT = 16384forfp8_grid(+7.3% with the consumer GEMM).Test Plan
python op_tests/test_flydsl_mega_mhc.py(gfx950; picked up byaiter_test.sh).post,no_post,identity_pre, allout_dtypes,T = 0..16400; ue8m0 outputs bit-exact against the torch rule on the kernel's own norm, every finish path.black,ruff.Test Result
Correctness:
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.bf16ISA byte-identical across the ue8m0 change; no spills or scratch.Performance vs SGLang v0.5.21 (kernel device time, both measured the same way):
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):T.T = 4160dips to 58%: 260 blocks on 256 CUs leave a near-empty second round; the split-K fallback recovers 179 → 124 µs.Tis latency-bound.Realistic Bandwidth Ceiling
At 16k rows a
postseam 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 atT = 4096, 419 MB: 71.6 µs).