Skip to content

perf(flydsl): MXFP4 fused-MoE stage2 optimization for EP prefill - #3117

Merged
inkcherry merged 15 commits into
ROCm:mainfrom
inkcherry:coopt/fusemoe-cleanup-v3-pr-ready
May 25, 2026
Merged

inkcherry merged 15 commits into
ROCm:mainfrom
inkcherry:coopt/fusemoe-cleanup-v3-pr-ready

Conversation

@inkcherry

@inkcherry inkcherry commented May 11, 2026

Copy link
Copy Markdown
Contributor

for InferenceMax case
cc @Duyi-Wang @GLRocks

Motivation

FlyDSL stage2 (down-projection) is on the hot path for DeepSeek-R1/V3 EP4 prefill on MI355X. The current _atomic_persist kernel doesn't overlap prologue VMEM with the K-loop, and flydsl_moe_stage2() redundantly memsets out even though moe_sorting_*_fwd already zeros moe_buf (~130us HBM bw per call).

Modifications

  • Add a production fp4×fp4 stage2 variant _t64x128x256_atomic_persist_async_w4_cumul3 (waves_per_eu=4, use_async_copy=True, cu_num_mul=3), plumb the new params through flydsl_moe_stage2compile_mixed_moe_gemm2.
  • Drop the defensive out.fill_(0) from flydsl_moe_stage2(); the standard fused_moe path zeros moe_buf via moe_buf_set_zero_kernel_2d. Caller contract documented in the wrapper.
  • Kernel-side passes (gated to the production tile + dtype): async X DMA in prologue, cu_num_mul=3 persistent grid expansion, asymmetric b_lo/b_hi split (3/4 + 1/4), LDS-staged sorted_weights, disable_xdl_arb_stall + carried s_setprio(1), scales-before-B in the K-loop, deferred lds_tid prologue + merged lds_tid/lds_tw guards, buffer-atomic fast path (prologue precomputes per-row byte offset + OOB mask), balanced contiguous M-tile dispatch.

Perf Tests

Hardware: single AMD MI355X.
Workload: test_fused_moe_ep4_mxfp4.pyM=49152, model_dim=7168, inter_dim=2048, E=64, topk=8, per_1x32 fp4x2×fp4x2 → bf16.
Metric: 5-run median × 3 reps, end-to-end fused_moe latency.

Source rep1 rep2 rep3 avg (us) Δ
upstream main 3709.9 3715.1 3710.4 3711.8
this PR 3371.2 3377.0 3406.5 3384.9 −8.81%

Correctness: test_flydsl_moe_a4w4.py 12/12 PASS across tile_m∈{32,64} × token∈{16,64,256} × {stage2, e2e}.

please help on review @lalala-sh,This has improved the inferenceMax E2E scenario. Thanks!

Optimize the FlyDSL stage2 (down-projection) kernel for the production
fp4xfp4 EP4 DeepSeek prefill shape (tile_m=64, tile_n=128, tile_k=256,
M=49152) on MI355X.

## Summary of changes

aiter/ops/flydsl/moe_kernels.py
  - Add new wrapper params (waves_per_eu, use_async_copy, cu_num_mul) and
    plumb them through compile_flydsl_moe_stage2 -> compile_mixed_moe_gemm2.
  - Register a production fp4xfp4 stage2 variant
    `flydsl_moe2_afp4_wfp4_bf16_t64x128x256_atomic_persist_async_w4_cumul3`
    with waves_per_eu=4, use_async_copy=True, cu_num_mul=3.
  - Drop the defensive `out.fill_(0)` added in PR ROCm#2863: when accumulate=True
    the standard fused_moe path already zeros moe_buf via
    `moe_buf_set_zero_kernel_2d` inside `moe_sorting_*_fwd`. The extra
    fill is a ~token_num*model_dim HBM write (~130 us per call at MI355X
    HBM bw on this shape). Callers using flydsl_moe_stage2 directly with
    accumulate=True remain responsible for zeroing `out` (documented).

aiter/ops/flydsl/kernels/mixed_moe_gemm_2stage.py
  - compile_mixed_moe_gemm2 signature: add waves_per_eu, use_async_copy,
    cu_num_mul. Drop unused legacy b_nt, xcd_swizzle on this path.
  - cu_num_mul: multiplies persistent-mode grid_y (CU count) by an integer
    factor; cu_num_mul>1 launches more persistent CTAs that each cover
    fewer M tiles, increasing in-flight parallelism. cu_num_mul=3 is the
    EP4-prefill sweet spot; cu_num_mul=4 regresses ~2.4% on the same shape.
  - use_async_copy: enable raw_ptr_buffer_load_lds for the X (activation)
    tile in the prologue so the X DMA overlaps with B/scale VMEM. Gated
    to the production tile/dtype combo so other configs are unaffected.
  - Asymmetric b_lo/b_hi split (b_lo loads 3/4 of K, b_hi loads 1/4): the
    smaller b_hi burst frees VMEM issue slots right before s_setprio(1)
    and the K-loop; the extra ku of b_lo is issued in the prologue where
    it overlaps with the X DMA.
  - Stage sorted topk weights into LDS once per M-tile and prefetch via
    vec4 LDS reads in the cshuffle epilogue, eliminating per-mi VMEM
    dwordx4 loads from the MFMA-heavy hot loop.
  - Emit AMDGPU `disable_xdl_arb_stall` + `s_setprio(1)` around the
    K-loop on gfx950 so the high-prio K-loop drains MFMAs back-to-back.
  - Issue scales BEFORE B-VMEM in the steady-state K-loop so the
    extract-time stalls in the first MFMA chunk are reduced.
  - Defer the lds_tid prologue (sorted_idx/sorted_w VMEM dwords + LDS
    stores) to AFTER the long-latency X DMA + B/scale loads, hiding the
    lds_tid scalar load latency behind the dominant VMEM phase.
  - Merge the lds_tid + lds_tw prologue store guards into a single
    scf.IfOp so MLIR can co-schedule the two VMEM buffer_loads.
  - For the buffer-atomic path: fast-path the cshuffle row token-index
    load (move the AND + multiply + validity-mask into the lds_tid
    preload site where it overlaps with B/scale/X VMEM loads).
  - Balanced persistent tile distribution (`tD4balPersist`) and balanced
    LDS-write hint accounting on the async path (`dswr_async`).

aiter/fused_moe.py
  - `_flydsl_stage2_wrapper`: forward waves_per_eu, use_async_copy,
    cu_num_mul from the parsed kernel-name config to flydsl_moe_stage2.

## Performance

Benchmark harness:
  /home/mingzliu/sgl_opt/ep4_mxfp4_opt/coopt_loop/scripts/bench.sh
Workload:
  /home/mingzliu/sgl_opt/test_fused_moe_ep4_mxfp4.py
  (M=49152, model_dim=7168, inter_dim=2048, 64 local experts, topk=8,
   QuantType.per_1x32, fp4x2 weights & activations, bf16 out)
Metric: 5-run median latency, 3 reps, MI355X (smci355-ccs-aus-n06-21).

  baseline (kernel=_atomic_persist):                3438.8 us (3-rep avg)
  this PR  (kernel=_atomic_persist_async_w4_cumul3): 3173.6 us (3-rep avg)
  delta:                                             -265 us  (-7.71%)
@inkcherry
inkcherry requested review from a team and Copilot May 11, 2026 05:43
@github-actions

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:triton-300x Run an additional Triton test job on MI300X in PRs; main branch always runs both MI35X and MI300X
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 3117 --add-label <label>

Copilot AI left a comment

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

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

Pull request overview

Optimizes the FlyDSL MoE stage2 (down-projection) path for MXFP4×MXFP4 EP prefill by introducing a production-tuned persistent+async kernel variant and reducing redundant memory traffic, with new tunable compilation parameters plumbed through the fused MoE wrapper.

Changes:

  • Add a new “production” fp4×fp4 stage2 kernel config (*_persist_async_w4_cumul3) and plumb waves_per_eu, use_async_copy, and cu_num_mul through stage2 compilation/dispatch.
  • Remove the defensive out.fill_(0) in flydsl_moe_stage2() for atomic mode (caller must provide a zeroed buffer when passing out).
  • Substantial stage2 kernel scheduling/memory pipeline changes (async X DMA, persistent grid expansion, LDS staging, buffer-atomic fast path, etc.) to improve overlap and latency.

Reviewed changes

Copilot reviewed 3 out of 3 changed files in this pull request and generated 5 comments.

File Description
aiter/ops/flydsl/moe_kernels.py Registers the new stage2 kernel config and plumbs new stage2 compilation knobs; updates stage2 output buffer initialization behavior.
aiter/ops/flydsl/kernels/mixed_moe_gemm_2stage.py Implements the stage2 kernel-side optimizations and adds new compilation parameters (waves_per_eu, async copy, CU multiplier).
aiter/fused_moe.py Passes newly added stage2 tuning parameters from parsed kernel configs into flydsl_moe_stage2().

💡 Add Copilot custom instructions for smarter, more guided reviews. Learn how to get started.

Comment thread aiter/ops/flydsl/moe_kernels.py
Comment thread aiter/ops/flydsl/kernels/mixed_moe_gemm_2stage.py
Comment thread aiter/ops/flydsl/kernels/mixed_moe_gemm_2stage.py
Comment thread aiter/ops/flydsl/moe_kernels.py Outdated
Comment thread aiter/ops/flydsl/kernels/mixed_moe_gemm_2stage.py Outdated
@inkcherry
inkcherry force-pushed the coopt/fusemoe-cleanup-v3-pr-ready branch from 2f0cdab to 70fd554 Compare May 11, 2026 06:21
The strict flydsl in current aiter-main rejects bare `if X:` / `elif X:`
over Python-level booleans inside a kernel emit body -- it lowers them
into `scf.if` so any Python variable assigned in either branch is
trapped inside the scf.if scope and the outer reference raises
`NameError: name '<X>' is not defined`.

`compile_mixed_moe_gemm1` already wraps every such condition with
`const_expr(...)`. The stage2 emit body in the parent commit had ~98
bare-if conditions over compile-time identifiers (`accumulate`,
`_persistent`, `doweight_stage2`, `_b_split_enabled`, `is_f4_a`,
`use_async_copy`, `_r139_xdma_first`, `_r216_defer_tid`,
`_use_buf_atomic_pre`, ...).  Wrap them with `const_expr(...)` so the
condition is evaluated at Python trace time and the variable binding
escapes to the enclosing scope, matching the stage1 convention.

No semantic change: only adds `const_expr(...)` around if/elif heads
inside compile_mixed_moe_gemm2; DSL/runtime conditions (those whose
sub-expressions reference `arith.`, `buffer_ops.`, `rocdl.`, ...) are
left as bare `if` so they remain runtime `scf.if`s.

Repro (without this commit) on aiter-main runtime:

  File ".../mixed_moe_gemm_2stage.py", line 3096, in moe_gemm2
      shape_lds = fx.make_shape(tile_m, _eff_lds_stride)
  NameError: name '_eff_lds_stride' is not defined

After this commit the kernel compiles, the end-to-end fused_moe perf
is unchanged from the parent commit on the legacy flydsl, and is
restored on aiter-main flydsl.
@inkcherry
inkcherry force-pushed the coopt/fusemoe-cleanup-v3-pr-ready branch from 70fd554 to 83ae8ad Compare May 11, 2026 06:25
inkcherry added 2 commits May 11, 2026 09:04
Re-flow a handful of long if-conditions, error message strings, and
multi-arg call sites in `aiter/ops/flydsl/moe_kernels.py` and
`aiter/ops/flydsl/kernels/mixed_moe_gemm_2stage.py` so they pass the
`psf/black@stable` check in `.github/workflows/pre-checks.yaml`.

No semantic change -- pure whitespace/line-wrap.  `ruff check` on the
three changed files is also clean.
Comment thread aiter/ops/flydsl/moe_kernels.py Outdated
Comment thread aiter/ops/flydsl/kernels/mixed_moe_gemm_2stage.py
Comment thread aiter/ops/flydsl/kernels/mixed_moe_gemm_2stage.py
Comment thread aiter/fused_moe.py
@coderfeli

Copy link
Copy Markdown
Collaborator

Update to main to merge in the hotfix flydsl version.

Comment thread aiter/ops/flydsl/kernels/mixed_moe_gemm_2stage.py
lalala-sh
lalala-sh previously approved these changes May 14, 2026
Comment thread aiter/ops/flydsl/kernels/mixed_moe_gemm_2stage.py
Wire xcd_swizzle end-to-end through compile_mixed_moe_gemm2 so the
fp4xfp4 stage2 path actually consumes it (mirrors stage1):
  - signature: drop the no-op `del xcd_swizzle`, keep b_nt as no-op
  - module_name: append `_xcd{N}` tag when xcd_swizzle > 0
  - kernel body: remap (bx_persist, by_outer) by NUM_XCDS=8 with the
    same group-of-M scheme stage1 uses, so consecutive WGs hit the
    same XCD and share L2 on MI355X
  - _cache_tag: include xcd_swizzle so distinct values JIT separately

Also pick up the black reflow on the production-variants tuple in
moe_kernels.py.
inkcherry added 4 commits May 15, 2026 05:20
The 2stage compile path was mis-classifying fp8 activations: only
is_f4_a triggered the e8m0 microscale buffer-resource setup, so
fp8xfp4 per_1x32 quantized inputs fell through to a wrong/empty
scale-buffer path. This produced large output mismatches in
op_tests/test_moe_2stage.py -q 7 (logits_diff ~0.98).

Widen the gate to (is_f4_a or is_f8_a) so fp8a + fp4w (a8w4)
correctly takes the e8m0 microscale buffer-descriptor branch,
matching the pre-perf-commit behaviour.

Verified locally on MI355X:
  - q=7 (a8w4) all 100 configs now PASS, max logits_diff = 1.19e-3
    (same magnitude as stock baseline, vs ~0.98 before)
  - 42 xcd_swizzle=4 routed variants also pass
  - no change to fp4xfp4 / bf16xint4 / fp16 paths (gate only
    widens the scale-buffer branch for fp8a)
This PR added `waves_per_eu` to the stage2 `_cache_tag` tuple, but
aot/flydsl/moe.py::_precompile_to_cache still defaulted it to `3`
while compile_mixed_moe_gemm2 defaults to `None`. AOT wrote the .pkl
keyed on (..., 3, ...), runtime looked up (..., None, ...) -> AOT
cache miss on every legacy stage2 kernel (test_moe_2stage.py first
fails at t32x128x256_atomic, token=4, fp4xfp4 per_1x32).

Align AOT default with runtime default. Stage1 and the new
`_persist_async_w4_cumul3` variant are unaffected because their
base_params always pin waves_per_eu explicitly.

Signed-off-by: inkcherry <mingzhi.liu@amd.com>
lalala-sh
lalala-sh previously approved these changes May 21, 2026
inkcherry added 2 commits May 21, 2026 15:59
Resolve trivial conflict in aiter/ops/flydsl/moe_kernels.py with
upstream PR ROCm#3274 which independently removed `out.fill_(0)` from
flydsl_moe_stage2(). Keep this PR's NOTE comment documenting why
the fill is unsafe to omit (caller-zeroed `moe_buf` in the fused_moe
dispatch path; ~130us HBM write per call on the EP4 prefill shape).

No code/behaviour change vs main; both sides converge on the same
``alloc_fn = torch.zeros if accumulate else torch.empty`` allocation.

Signed-off-by: inkcherry <mingzhi.liu@amd.com>
@inkcherry
inkcherry merged commit 4f6f591 into ROCm:main May 25, 2026
27 of 30 checks passed
valarLip added a commit that referenced this pull request May 26, 2026
zhiding512 added a commit that referenced this pull request May 26, 2026
PR #3117 (perf(flydsl): MXFP4 fused-MoE stage2 optimization for EP
prefill) was reverted upstream in #3344, but its changes leaked into
this branch when the pad-runtime work was copied wholesale from
zhimding/moe_pad_runtime_0526 (which was based on a pre-revert local
main). Strip out:

- _register_production_variants_stage2 + its hand-tuned EP4 variant
- waves_per_eu/use_async_copy/cu_num_mul plumbing in stage2 compile
- stage2 kernel-level perf changes

Keep only the pad-runtime delta: stage1/stage2 compile_* drop
model_dim_pad/inter_dim_pad as compile-time params; stage2 kernel
takes them as runtime SSA via i32_*_pad_in; K-tail unroll no longer
specializes on pad (weight pad region is zero -- full k_unroll is
correct and equally fast).

Signed-off-by: root <zhimding@amd.com>
inkcherry added a commit that referenced this pull request Jun 11, 2026
* Re-introduce #3117 MXFP4 fused-MoE stage2 EP-prefill opt with K-pad NaN fix

Reverts upstream's revert (#3344) of #3117 and lands the surgical stage2
K-direction padding tail-skip fix on top, so the EP-prefill MXFP4 2-stage
optimization is restored without the GSM8K accuracy collapse (0.97 -> 0.00
under AITER_BF16_FP8_MOE_BOUND=0 + ATOM_MOE_GU_ITLV=1 + fp8 KV).

- mixed_moe_gemm_2stage.py: taken from the validated fix (ebe7f55) =
  #3117 stage2 rewrite + ku_count tail-skip + tile2_pad grid correction.
- fused_moe.py: keep main's _get_padding_for_flydsl helper; #3117 caller
  path re-applied.
- moe_kernels.py / aot/flydsl/moe.py / dsv3_fp4_tuned_fmoe.csv: #3117
  re-applied (auto-merged with current main).

* port #3117 stage2 MoE kernel to new FlyDSL buffer/ptr API

Adapt mixed_moe_gemm_2stage.py (the re-introduced #3117 MXFP4 fused-MoE
stage2 EP-prefill optimization) to the FlyDSL API migration landed in
#3373 ("use fx.ptr for kernels not using layout"), so it builds and runs
on the current aiter/main + nightly FlyDSL (0.1.9.dev599).

Without this, kernel compile fails MLIR verification because the new
FlyDSL passes pointer kernel args (!fly.ptr) where the old
buffer_ops.create_buffer_resource(...) expected a memref and emitted
fly.extract_aligned_pointer_as_index on a pointer.

Migration (mirrors #3373 applied to this same file):
- kernel/launch args: fx.Tensor -> fx.Pointer
- add _ptr_buffer_resource() helper using fx.ptrtoint +
  create_buffer_resource_from_addr in both compile_mixed_moe_gemm1/2
- add w_nbytes/bias_nbytes byte-size computations for the resources that
  previously relied on max_size=False
- output base address: extract_aligned_pointer_as_index(arg_out) ->
  arith.index_cast(i64, fx.ptrtoint(arg_out))

#3117's perf logic and the K-pad NaN tail-skip fix are untouched.

Validated E2E on DeepSeek-V4-Pro (ATOM, tp8, fp8 KV) under the regression
trigger env AITER_BF16_FP8_MOE_BOUND=0 ATOM_MOE_GU_ITLV=1:
GSM8K exact_match 0.9733 vs nightly baseline 0.9533 (equal within noise);
previously this config produced garbage output (0.00).

* test(moe_2stage): add deterministic NaN guard for #3117 stage2 regression

The MXFP4 fused-MoE stage2 EP-prefill regression from #3117 surfaces as
NaN in the fused_moe output (uninitialized K-padding tail-tile lanes).
checkAllclose's err/logits_diff can be masked by atomic-reduction noise,
so add an explicit out2_ck.isnan() check: assert in strict_accuracy mode
(used by the tuned-config sweep) and warn otherwise. This deterministically
fails on the unfixed #3117 kernel and passes on the K-pad fix.

* style: black-format test_moe_2stage NaN guard

psf/black@stable wrapped the logging.error(...) call added in the
previous commit (line exceeded 88 cols). No logic change.

* fix(moe2): use Python-int K offset for stage2 scale-shift (fix AOT std::bad_cast)

The steady-state K-loop in compile_mixed_moe_gemm2 passed the MLIR runtime
value `next_k1` (= arith.index(k_iv_py) + tile_k) into `_k_shift_bits` /
`_k_base`, but those helpers build a compile-time
`arith.constant(k_shift_bits, type=i32)` and therefore require a Python int
-- mirroring the prologue's `_k_shift_bits(0)`, the next_k2 step's
`_k_shift_bits(next_k2_py)`, and the tail's `_k_shift_bits(k_tail1_py)`.

Feeding an MLIR Value into arith.constant raises `Invalid attribute value
for the key "value" ... (std::bad_cast)`, which broke FlyDSL MOE AOT
pre-compilation in build_aiter_wheels for the 32 flydsl_moe2_afp4_wfp4_*
(a4w4, tile_k=128) stage2 variants. Only tile_k=128 hits it: tile_k=256 has
pack_K >= _scale_pack_k, so `_k_shift_bits` returns 0 before reaching the
constant (hence the a8w4 path was unaffected).

Add `next_k1_py = k_iv_py + tile_k` and use it for both helpers. Verified
all 32 previously-failing a4w4 variants now compile and the a8w4 DSv4 path
still compiles via the MOE AOT.

* [flydsl] Fix a4w4 fused-MoE stage1 grid padding (tile_k_stage2)

compile_mixed_moe_gemm1 hard-coded `tile_k_stage2 = 256` when deriving
`tile2_pad` and the stage1 grid-X extent. That value is only correct for
the fp8-activation (a8w4) kernels whose stage1 tile_k = 512, where
`tile_k // 2 == 256`. The double-MXFP4 (a4w4) kernels use stage1
tile_k = 256, so the correct value is `tile_k // 2 == 128`. The
hard-coded 256 (plus the dropped outer `% tile_k_stage2` wrap) produced a
wrong grid-X for every a4w4 shape and a systematic ~0.02 logits_diff,
failing the strict shard-7 accuracy gate.

Restore `tile_k_stage2 = tile_k // 2` and the modulo wrap (the pre-#3117
behaviour). a8w4 kernels (tile_k = 512) are unaffected since 512 // 2
still equals 256; the #3117 stage2 (compile_mixed_moe_gemm2) optimisation
is untouched.

Verified on MI355X (gfx950), gptoss 3072/3072 E128 topk4, strict gate:
  a4w4 token=256             : 0.0206 (FAIL) -> 6.2e-06 (PASS, == main)
  a8w4 token=512/1024/2048   : 0.0048 (unchanged, bit-identical to main)

* fix(flydsl/moe2): restore #3476 inter_dim-128-align scale padding

The #3117 stage2 re-introduction wholesale-replaced
mixed_moe_gemm_2stage.py with a copy predating #3476 ("Support moe
inter_dim align with 128"), silently dropping #3476's microscale
padding. The e8m0 microscale buffers were then sized/strided with the
raw inter_dim, while the host e8m0_shuffle pads scale group-N up to a
multiple of 8 (= inter_dim rounded to the next 256). For inter_dim that
is 128- but not 256-aligned (e.g. the minimax_m25 tuned row:
model_dim=3072 inter_dim=384 E=256 topk=8, a4w4 persist_sbm64), the
kernel read OOB/mismatched scales -> garbage e8m0 exponents -> 100%
inf/NaN stage2 output, tripping the strict NaN guard in
test_moe_2stage.py (CI: MI35X Standard Tests shard 2).

Re-apply #3476's padding in both stages:
- gemm1: _sorted_scale_cols rounded up to a multiple of 8.
- gemm2: add scale_k_padded / scale_kblk_padded (inter_dim rounded to
  next 256) and use them for the A/B scale preshuffle layout (c_k_orig)
  and the A2/W microscale buffer K/32 strides (kblk / kblk_w).

No-op when inter_dim is 256-aligned (e.g. the a8w4 DSv4 path), so it
does not affect the previously-validated configs.

Verified on MI355X: the minimax a4w4 persist_sbm64 config goes from
100% NaN to logits_diff=6.2e-4 (< 0.01 strict gate); a 256-aligned
a4w4 row (inter_dim=256) still passes.
lalala-sh added a commit that referenced this pull request Jun 13, 2026
* Tune gfx1151 MHA forward default tile config (#3560)

Retune the fwd.default branch of the gfx1151 MHA Triton config. The
previous values were inherited from a large-tile (CDNA-style) layout that
suits MFMA-based architectures. gfx1151 (RDNA3.5) has no MFMA and a smaller
LDS/occupancy budget, so a smaller tile with software pipelining and lower
occupancy pressure is consistently faster.

Found via a coarse-then-refine sweep over BLOCK_M/BLOCK_N/num_warps then
waves_per_eu/num_stages/PRELOAD_V, scored by per-shape-normalized geomean
across a representative head-dim basket (d=64 with sink+sliding-window,
d=88, d=128).

Changes:
- fwd.default: BLOCK_M 128->64, BLOCK_N 64->32, num_warps 8->4,
  waves_per_eu 2->1, num_stages 1->2 (PRELOAD_V/num_ctas unchanged).
- fwd.pe and the backward branches are left unchanged.

Measured on forward, batch=1, bf16, causal, seq 1K-8K, both bshd and thd
layouts, min of two runs: median +8.3% (mean +11.6%) vs the prior
gfx1151 default across the dense models, up to +37% on the d=64
sink+sliding-window shape, with no regressions.

* Add hip mhc_fused_post_pre (#3623)

* Add mhc_fused_post_pre_gemm_sqrsum

* Fix accuracy

* optimize

* add tile_m=32

* add tile mnk to dispatch config

* add gfx942 80 cu config

* add dispatch logic : when m > 64 use unfuse

* optimize mhc_pre_big_fuse_kernel&mhc_pre_big_fuse_rmsnorm_kernel

* mhc_fused_post_pre_gemm_sqrsum : in-kernel hc_mult reduce

* update dispatch

* Fix 192x128 kernel deterministic bug (#3633)

* [module_custom] refactor (#3625)

* [module_causal_conv1d_update] refactor hip kernel (#3595)

* [module_causal_conv1d_update] refactor hip kernel

* restore comment

* [FlyDSL MoE] Add no combine feature (#3408)

* [OPUS] gfx942 a16w16 bf16 GEMM pipeline family for DSV4 (#3594)

* feat(opus): add gfx942 EM3EN4 LDS1 PGR2 splitK path

* feat(opus): add gfx942 WKC paths and reduce cleanup

* perf(opus): add gfx942 bf16 workspace splitK path

* perf(config): refresh gfx942 OPUS bf16 model configs

* style(opus): format splitK reduce TU check

---------

Co-authored-by: yifehuan <yifehuan@local>

* CI: map ATOM MI350X runner label (#3647)

* CI: map ATOM MI350X runner label

* CI: route ATOM MI350X label to DO runner

* CI: limit ATOM and SGLang test concurrency

* add env var for kernel arg preload (#3649)

* add env var for kernel arg preload

* fix

* fix name

* Refine flydsl gemm config selection code (#3608)

* Update tuned_gemm.py

* Update tuned_gemm.py

* Fix opus gemm aiter check (#3622)

* 1. add g_aiter_can_throw = true to throw runtime_error not abort; 2. refactor rtol/atol in gemm_a16w16_tune

* update README.md

* ds dummy pass

* mha_native: native HIP D64 BF16 split-K forward backend for flash_attn_func (#3581)

* mha_native: vendor fmha_native D64 device headers (verbatim @ f7daf17)

* mha_native: split-K producer/combine entries with intra-TU launch wrappers

* mha_native: orchestrator (scratch, scale fold, stream, producer+combine)

* mha_native: pybind module + rocm_ops macro

* mha_native: register module_mha_fwd_native_splitkv JIT config

* mha_native: python stub + gen_fake; kernel builds and runs

* mha_native: capability gate, split heuristic, dispatch branch

* mha_native: thread num_splits through flash_attn_func autograd chain

* mha_native: add -ns/--num_splits to test_mha.py call chain

* mha_native: add -mllvm -enable-post-misched=1 for ISA parity with fmha_native

* mha_native: clean up code comments

* mha_native: drop stale source-repo comment references

Comments in the vendored device headers referred to "the four existing
entries", "four call sites" and "[_varlen]" entry files that exist in the
upstream source but not on this branch (only the msk{0,1}_split producers
and combine ship here; fmha_fwd_d64_device is always instantiated with
IsSplit=true, IsVarlen=false). Reword to describe this branch's actual
entry set. Comments only; no code lines changed (ISA parity preserved).

Co-Authored-By: Claude Opus 4 (1M context) <noreply@anthropic.com>

* mha_native: replace hardcoded split-KV heuristic with tuned occupancy model

Port the two-regime occupancy heuristic tuned on 100 measured D64 shapes,
replacing the single hardcoded special case. G==0 falls back to the CK
non-split-KV kernel. Thread seqlen_q through and detect CU count dynamically
via get_cu_num() instead of assuming 304.

Co-Authored-By: Claude Opus 4 (1M context) <noreply@anthropic.com>

* mha_native: correct stale ns==1 fall-through comment

The heuristic now also returns 0 (CK fallback), so the dispatch falls
through for ns <= 1, not just ns == 1. Comment only.

Co-Authored-By: Claude Opus 4 (1M context) <noreply@anthropic.com>

* mha_native: apply black formatting to mha.py

Co-Authored-By: Claude Opus 4 (1M context) <noreply@anthropic.com>

* mha_native: validate inputs at native splitkv entrypoint

mha_fwd_native_splitkv is a public aiter:: symbol exposed via pybind, so
C++/direct callers bypass the Python-side can_impl_fmha_native gating.
Validate k/v dtype (was q-only), 4-D rank, q/k/v last-dim contiguity, and
Hq % Hk == 0 (Hk > Hq would divide by zero in device GQA grouping).

Co-Authored-By: Claude Opus 4 (1M context) <noreply@anthropic.com>

* mha_native: validate user-supplied out tensor at native splitkv entrypoint

out_opt is written via reinterpret_cast to bf16 using its own strides over
a grid covering B*Hq*Sq*D, so a wrong dtype/device/shape silently corrupts
memory or writes out of bounds. Validate bf16 dtype, same device as q, and
(B,Sq,Hq,D) shape in addition to the existing last-dim contiguity check.

Co-Authored-By: Claude Opus 4 (1M context) <noreply@anthropic.com>

* mha_native: reject negative num_splits and document its semantics

- Reject num_splits < 0 explicitly; previously negatives silently fell
  through to the heuristic path like 0, masking misconfiguration.
- Document num_splits in the flash_attn_func docstring (0=auto/heuristic,
  1=disable split-K, >=2 forces native split-K when applicable).
- Fix the test CLI --num_splits help text, which claimed ">=1 forces
  native" while the dispatch only routes to native when num_splits >= 2.

Co-Authored-By: Claude Opus 4 (1M context) <noreply@anthropic.com>

* mha_native: reject finite-0 window and sink in native gate

can_impl_fmha_native used `not swa` (swa = left>0 or right>0), which let a
finite 0 window slip through -- e.g. window_size=(-1, 0) is semantically a
causal mask, but with causal=False the native dispatch ran the unmasked
kernel, diverging from CK by ~3.5 max-abs (verified). Require the exact
no-window sentinel (left==-1 and right==-1) and sink_size==0 so any
window/sink restriction falls back to CK/ASM. Native's two supported modes
(full / full-causal) are unchanged; verified no regression.

Co-Authored-By: Claude Opus 4 (1M context) <noreply@anthropic.com>

* mha_native: bound Q/K/V buffer SRDs to fix OOB read on uneven seqlen_k

The Q/K/V buffer resource descriptors used num_records=0xFFFFFFFF, which
disables the hardware bounds check. When seqlen_k % kN0 != 0 the K/V tile
loop walks a full kN0-wide tile, so padding rows (row >= seqlen_k) read
past the tensor into adjacent/freed memory. Those reads are masked to 0 in
the score, but P(=0)*V(=garbage) computes 0*NaN = NaN in GEMM1, which
poisons O_acc -> nondeterministic NaN output on uneven-seqlen_k cases.

Bind each SRD to the valid byte extent of its (b,h) region so OOB reads
return 0 (the hardware's defined behavior) instead of garbage, restoring
the read-padding-then-mask invariant.

Co-Authored-By: Claude Opus 4 (1M context) <noreply@anthropic.com>

* mha_native: seed rmax below any realizable raw score

The online-softmax running max was seeded at a finite -5000, but
large-magnitude inputs produce raw Q.K^T scores past -5000 for a causal
row that sees a single strongly-negative key. The seed then won the
running max, the lone score exp2-underflowed to 0, and the row was
wrongly treated as fully masked (O=0, LSE=-inf). Seed at -1e30 instead:
far below any realizable raw score yet finite, so a real score always
wins while genuinely masked rows still collapse cleanly to O=0/LSE=-inf
without producing NaN.

Co-Authored-By: Claude Opus 4 (1M context) <noreply@anthropic.com>

---------

Co-authored-by: Claude Opus 4 (1M context) <noreply@anthropic.com>
Co-authored-by: Jim Guo <jim.guo@amd.com>

* [Feat] Support FP4 gather_kv_b_proj (#3597)

* Support FP4 gather_kv_b_proj

Add an FP4 per-1x32 MXFP4 path for gather_kv_b_proj so cached MLA KV expansion can run with raw and preshuffled FP4 weights.

* refactor FP4 gather_kv_b_proj path

Merge the FP4 MXFP4 gather path into the existing Triton gather kernel and add coverage for raw and preshuffled FP4 weights.

* Fix FP4 gather_kv_b_proj launch grid

Avoid sizing the FP4 chunk-parallel launch from kv_indices capacity and flatten the chunk dimension into a 1D grid so serving metadata cannot exceed the AMD grid-y launch limit. Add coverage for oversized kv_indices buffers that only have a smaller valid range.

* Optimize FP4 gather_kv_b_proj tiling

Tune the FP4 gather path to avoid empty chunk work and use layout-specific dot_scaled granularity so raw and preshuffled MXFP4 layouts stay performant in serving shapes.

* [Gluon][gfx1250] gemm a16w16 cleanup (#3646)

Add Gluon BF16 GEMM for gfx1250

Co-authored-by: ahmed-bsod <Muhammad.Ahmed@amd.com>
Co-authored-by: Vinayak Gokhale <Vinayak.Gokhale@amd.com>
Co-authored-by: Omar Muhammad <omar.muhammad@amd.com>
Co-authored-by: Alexander Weinrauch <alexander.weinrauch@amd.com>
Co-authored-by: Lukasz Burzawa <lukasz.burzawa@amd.com>

* [Gluon] add torch compile guard + tdm descriptor fix in routing.py (#3530)

* test: skip pa_decode_bf16_asm off gfx1250 (#3660)

* [Triton] [Gluon] [GFX12] UA3D update config (#3612)

* update

* config

* config

* config

* config

* config

* config, fix

* format

* reorder

* fix routing

* fix routing

* update repr

* update config

* change qk_factor back to multiply at for loop

---------

Co-authored-by: Vinayak Gokhale <Vinayak.Gokhale@amd.com>
Co-authored-by: Mehmet Cagri Kaymak <mehmet.kaymak@amd.com>

* update

* fea: reduce_scatter support all dim (#3464)

* fea: reduce_scatter support all dim

Signed-off-by: TennyWang1223 <root@hjbog-srdc-24.amd.com>

* [fix]: support dim!=0 case in torch distribute

Signed-off-by: TennyWang1223 <root@hjbog-srdc-24.amd.com>

---------

Signed-off-by: TennyWang1223 <root@hjbog-srdc-24.amd.com>
Co-authored-by: TennyWang1223 <root@hjbog-srdc-24.amd.com>

* Fix assert in triton fused_kv_cache (#3601)

* verify pass

* Add MiniMax M2.5 FMoE tuned config updates (#3644)

Merge newly tuned MiniMax FMoE entries while keeping the best latency per shape and avoiding duplicate runtime config keys.

Signed-off-by: Aakif Nawaz <aakif.nawaz@amd.com>
Co-authored-by: yzhou103 <Ying.Zhou2@amd.com>

* update

* Fix HK MLA decode fwd: per-batch output bounds check + reduce-time per-tile split cap (#3391)

Fix an page fault in HK MLA decode that surfaces under ragged per-batch qseqlen. The fault was reported in kn_mla_reduce_v1<512,16,4> but the root cause is in the HK fwd kernel's output store, which silently corrupts adjacent VRAM (including the reduce metadata tensors).

Two related changes:

1. Per-batch bounded SRD in HK fwd `OManager*::output_to_vram`. The HK m16x4 / m16x8 OutputFinal path always writes a full kBlockM (=64) rows of bf16 per work item, but when a trailing batch has `qo_len < max_seqlen_q`, the surplus rows overshoot final_output and stomp whatever PyTorch placed adjacent in VRAM (in our case work_indptr / work_info_set / reduce_partial_map, which the reduce kernel then dereferences with garbage offsets). The asm path doesn't have this bug because it builds a per-batch bounded SRD; HK was using `num_records = 0xFFFFFFFF`. Add a bool `kCheckOOB` template parameter and qo_end runtime parameter to all four `OManager{16,32}bits{V1,V2}::output_to_vram` methods.

2. Reduce-time per-tile split cap plumbing. The reduce kernel previously sized its LDS state via `params.max_splits = multiProcessorCount`, which under-allocates when the metadata kernel was launched at `cu_num * occupancy` (HK m16x4 runs at `occupancy=2`). Plumb `num_kv_splits` through `mla_reduce_v1` (host wrapper, pybind, header) so `params.max_splits = max(multiProcessorCount, num_kv_splits)`. Factor the `cu_num * occupancy` computation into two new helpers in `aiter/ops/attention.py`.

* [fix](gemm): fix hang issue (#3664)

Co-authored-by: perzhang <perzhang@amd.com>

* update tunner adn token num

* add gptoss and ds gemm config for gfx1250 (#3676)

Co-authored-by: Vinayak Gokhale <Vinayak.Gokhale@amd.com>

* readd #3117, EP prefill optimization 

* Re-introduce #3117 MXFP4 fused-MoE stage2 EP-prefill opt with K-pad NaN fix

Reverts upstream's revert (#3344) of #3117 and lands the surgical stage2
K-direction padding tail-skip fix on top, so the EP-prefill MXFP4 2-stage
optimization is restored without the GSM8K accuracy collapse (0.97 -> 0.00
under AITER_BF16_FP8_MOE_BOUND=0 + ATOM_MOE_GU_ITLV=1 + fp8 KV).

- mixed_moe_gemm_2stage.py: taken from the validated fix (ebe7f55) =
  #3117 stage2 rewrite + ku_count tail-skip + tile2_pad grid correction.
- fused_moe.py: keep main's _get_padding_for_flydsl helper; #3117 caller
  path re-applied.
- moe_kernels.py / aot/flydsl/moe.py / dsv3_fp4_tuned_fmoe.csv: #3117
  re-applied (auto-merged with current main).

* port #3117 stage2 MoE kernel to new FlyDSL buffer/ptr API

Adapt mixed_moe_gemm_2stage.py (the re-introduced #3117 MXFP4 fused-MoE
stage2 EP-prefill optimization) to the FlyDSL API migration landed in
#3373 ("use fx.ptr for kernels not using layout"), so it builds and runs
on the current aiter/main + nightly FlyDSL (0.1.9.dev599).

Without this, kernel compile fails MLIR verification because the new
FlyDSL passes pointer kernel args (!fly.ptr) where the old
buffer_ops.create_buffer_resource(...) expected a memref and emitted
fly.extract_aligned_pointer_as_index on a pointer.

Migration (mirrors #3373 applied to this same file):
- kernel/launch args: fx.Tensor -> fx.Pointer
- add _ptr_buffer_resource() helper using fx.ptrtoint +
  create_buffer_resource_from_addr in both compile_mixed_moe_gemm1/2
- add w_nbytes/bias_nbytes byte-size computations for the resources that
  previously relied on max_size=False
- output base address: extract_aligned_pointer_as_index(arg_out) ->
  arith.index_cast(i64, fx.ptrtoint(arg_out))

#3117's perf logic and the K-pad NaN tail-skip fix are untouched.

Validated E2E on DeepSeek-V4-Pro (ATOM, tp8, fp8 KV) under the regression
trigger env AITER_BF16_FP8_MOE_BOUND=0 ATOM_MOE_GU_ITLV=1:
GSM8K exact_match 0.9733 vs nightly baseline 0.9533 (equal within noise);
previously this config produced garbage output (0.00).

* test(moe_2stage): add deterministic NaN guard for #3117 stage2 regression

The MXFP4 fused-MoE stage2 EP-prefill regression from #3117 surfaces as
NaN in the fused_moe output (uninitialized K-padding tail-tile lanes).
checkAllclose's err/logits_diff can be masked by atomic-reduction noise,
so add an explicit out2_ck.isnan() check: assert in strict_accuracy mode
(used by the tuned-config sweep) and warn otherwise. This deterministically
fails on the unfixed #3117 kernel and passes on the K-pad fix.

* style: black-format test_moe_2stage NaN guard

psf/black@stable wrapped the logging.error(...) call added in the
previous commit (line exceeded 88 cols). No logic change.

* fix(moe2): use Python-int K offset for stage2 scale-shift (fix AOT std::bad_cast)

The steady-state K-loop in compile_mixed_moe_gemm2 passed the MLIR runtime
value `next_k1` (= arith.index(k_iv_py) + tile_k) into `_k_shift_bits` /
`_k_base`, but those helpers build a compile-time
`arith.constant(k_shift_bits, type=i32)` and therefore require a Python int
-- mirroring the prologue's `_k_shift_bits(0)`, the next_k2 step's
`_k_shift_bits(next_k2_py)`, and the tail's `_k_shift_bits(k_tail1_py)`.

Feeding an MLIR Value into arith.constant raises `Invalid attribute value
for the key "value" ... (std::bad_cast)`, which broke FlyDSL MOE AOT
pre-compilation in build_aiter_wheels for the 32 flydsl_moe2_afp4_wfp4_*
(a4w4, tile_k=128) stage2 variants. Only tile_k=128 hits it: tile_k=256 has
pack_K >= _scale_pack_k, so `_k_shift_bits` returns 0 before reaching the
constant (hence the a8w4 path was unaffected).

Add `next_k1_py = k_iv_py + tile_k` and use it for both helpers. Verified
all 32 previously-failing a4w4 variants now compile and the a8w4 DSv4 path
still compiles via the MOE AOT.

* [flydsl] Fix a4w4 fused-MoE stage1 grid padding (tile_k_stage2)

compile_mixed_moe_gemm1 hard-coded `tile_k_stage2 = 256` when deriving
`tile2_pad` and the stage1 grid-X extent. That value is only correct for
the fp8-activation (a8w4) kernels whose stage1 tile_k = 512, where
`tile_k // 2 == 256`. The double-MXFP4 (a4w4) kernels use stage1
tile_k = 256, so the correct value is `tile_k // 2 == 128`. The
hard-coded 256 (plus the dropped outer `% tile_k_stage2` wrap) produced a
wrong grid-X for every a4w4 shape and a systematic ~0.02 logits_diff,
failing the strict shard-7 accuracy gate.

Restore `tile_k_stage2 = tile_k // 2` and the modulo wrap (the pre-#3117
behaviour). a8w4 kernels (tile_k = 512) are unaffected since 512 // 2
still equals 256; the #3117 stage2 (compile_mixed_moe_gemm2) optimisation
is untouched.

Verified on MI355X (gfx950), gptoss 3072/3072 E128 topk4, strict gate:
  a4w4 token=256             : 0.0206 (FAIL) -> 6.2e-06 (PASS, == main)
  a8w4 token=512/1024/2048   : 0.0048 (unchanged, bit-identical to main)

* fix(flydsl/moe2): restore #3476 inter_dim-128-align scale padding

The #3117 stage2 re-introduction wholesale-replaced
mixed_moe_gemm_2stage.py with a copy predating #3476 ("Support moe
inter_dim align with 128"), silently dropping #3476's microscale
padding. The e8m0 microscale buffers were then sized/strided with the
raw inter_dim, while the host e8m0_shuffle pads scale group-N up to a
multiple of 8 (= inter_dim rounded to the next 256). For inter_dim that
is 128- but not 256-aligned (e.g. the minimax_m25 tuned row:
model_dim=3072 inter_dim=384 E=256 topk=8, a4w4 persist_sbm64), the
kernel read OOB/mismatched scales -> garbage e8m0 exponents -> 100%
inf/NaN stage2 output, tripping the strict NaN guard in
test_moe_2stage.py (CI: MI35X Standard Tests shard 2).

Re-apply #3476's padding in both stages:
- gemm1: _sorted_scale_cols rounded up to a multiple of 8.
- gemm2: add scale_k_padded / scale_kblk_padded (inter_dim rounded to
  next 256) and use them for the A/B scale preshuffle layout (c_k_orig)
  and the A2/W microscale buffer K/32 strides (kblk / kblk_w).

No-op when inter_dim is 256-aligned (e.g. the a8w4 DSv4 path), so it
does not affect the previously-validated configs.

Verified on MI355X: the minimax a4w4 persist_sbm64 config goes from
100% NaN to logits_diff=6.2e-4 (< 0.01 strict gate); a 256-aligned
a4w4 row (inter_dim=256) still passes.

* [Triton] Sage MXFP4 return LSE (#3349)

* return lse from sage mxfp4, guarded by paramter return_lse (default false)

* black

---------

Co-authored-by: Juuso Korhonen <40278371+juuso-oskari@users.noreply.github.com>

* Add and tune fused GEMM A8W8 blockscale A16W16 benchmark (#3568)

* Add and tune fused GEMM A8W8 blockscale A16W16 benchmark

* fix import sys changeand main()

* gfx950 MoE A8W4: tuned entries for gpt-oss shapes + fallback hardening (#3580)

* gfx950: add A8W4 dispatch entries for 15 shapes the proxy-fallback mistunes

get_kernel_config_triton() falls back to an arch-heuristic when a (bm, N, K)
lookup misses gfx950-A8W4.json. For 15 (bm, N, K) tuples exercised by
gpt-oss-120b W4A8 at TP={1,2,4}, the heuristic picks BLOCK_SIZE_N in {256,512}
with num_stages=1, while aiter 0.1.13 (via gfx950-MOE-MX_FP4_A8.json) and
direct micro-tuning both prefer BLOCK_SIZE_N=128, num_stages=2 for the same
shapes.

Adds entries for bm{32,64,128} x N{1536,3072,6144} x K{768,1536,3072} (skipping
those already present), all tuned to:
  BLOCK_SIZE_K=256, BLOCK_SIZE_N=128, matrix_instr_nonkdim=16,
  num_stages=2, num_warps=4, waves_per_eu=0

Config probe on gpt-oss-120b W4A8 (TP=1/CONC=32/ISL=1024/OSL=1024) shows
106 of 210 shared (M, N, K, bm) shapes pick a different config between
aiter 0.1.13 and aiter HEAD; all 106 divergences trace back to these 15
missing entries.

* gfx950 A8W4: fix proxy BLOCK_K hazard, pick num_stages by LDS

  - Proxy fallback now skips BLOCK_K<256 (CDNA4 unswizzle won't compile).
  - gfx950 heuristic uses pick_gemm_num_stages instead of hardcoded ns=1.

* remove entry bm16_n2880_k360, as it's only 3.5% perf difference

* Fix CI LDS OOM in fallback + CDNA4 swizzle BK<256 crash

* address review: recove deleted json configs entry to avoid perf regression and use_async_padding=True for async-copy lowering on gfx950

* address review: add swizzle_mx_scale into get_kernel_config_triton and short-circuit the BLOCK_K>=256 check unless the caller is actually using CDNA4_SCALE

* [Triton] Add New Features and Performance Improvement for GMM Kernel (#3407)

* Support new transposed layouts

Transposed RHS in GMM:
> shape = (G, N, K) and stride = (K*N, K, 1)

Transposed LHS in TGMM:
> shape = (M, K) and stride = (K, 1)

* Tune `gfx950` GMM config

* Support `int64` data type for group sizes

* Add `grid_dim` override option to persistent kernel wrappers

The affected kernel wrappers are GMM and persistent TGMM.

* Disable expensive GMM assertion logic that uses GPU

* Implement work stealing GMM

* Add unit tests for the new features

The new features are:

* alternate transposed layouts (transposed RHS in GMM and transposed LHS
  in TGMM)
* `int64` data type for group sizes tensor
* grid dimension override for persistent kernels (GMM, PTGMM)
* work stealing GMM

* Add new production shape to GMM benchmark script

* [Review] Remove wrong assertion

Initial `tile == program_id` may be `>= total_tiles` when
`GRID_DIM > total_tiles`. Cheap `_gmm_grid` path returns
`num_programs = grid_dim`without clamping.

* [Review] Collapse `_gmm_grid` into one cheap path and one expensive path

* [Review] Defensive dict copy when overriding `grid_dim`

* [Review] Cache tile counter of work stealing GMM

* [triton-mha] hint head-stride div-by-8 for vectorized global load (#3424)

For the packed `[seq, heads, dim]` layout used by varlen prefill, the
head-axis stride equals `head_dim`. When `head_dim` is a multiple of 8
but not 16 (e.g. 72), Triton's integer-arg auto-specialization does not
attach `tt.divisibility = 8` to `stride_*h` (its threshold is 16), so
AxisInfo treats the K/V global load as 2-byte aligned and Coalesce
emits scalar `buffer_load_u16` instead of vectorized `buffer_load_b128`.

Add a `HEAD_STRIDE_ALIGNED_8` constexpr to `_attn_fwd` and apply
`tl.multiple_of(off_h_{q,k} * stride_{q,k,v}h, 8)` to the head-axis
integer offset when the caller sets it. AddPtr propagates this through
to the load pointer, so AxisInfo computes a 16-byte alignment and the
load coalesces.

The wrapper checks `stride_*h % 8 == 0` against the actual runtime
strides (not against `head_dim`), so the hint stays sound for
non-contiguous Q/K/V views where `stride_*h != head_dim`. The constexpr
defaults to `False`, so external callers of `_attn_fwd` are unaffected
unless they opt in.

Mirrors the equivalent hint added to the `flash_attn_triton_amd`
(`dao_ai`) prefill kernel; together with the gfx1151 tuning config this
closes the gap with the `dao_ai` impl on Strix Halo.

* Enable FP8 varlen MHA tests with latest Triton (#3643)

* Drop the loop carried percentage (#3661)

the code gen is the same with or without this flag with the LLVM branch (amd/dev/aukerbow/CoExecScheduler-staging)

* [Triton][CDNA4] Optimize gluon blockscale a8w8 gemm kernel (#3307)

* gluon blockscale: enable gluon impl in tests + bench correctness check

* test_gemm_a8w8_blockscale: enable the "gluon" parametrize entry and
  add small-K shapes (K in {128, 192, 256, 320}) that exercise the
  wind-down's num_k_iter guards.
* bench_gemm_a8w8_blockscale: add a -test flag that runs each
  benchmarked shape against a torch reference via checkAllclose.

* gluon blockscale: mfma_scaled pipeline

Reimplement the gluon a8w8 blockscale kernel around
gl.amd.cdna4.mfma_scaled with an explicit async-copy / LDS multi-buffer
pipeline.

* Split the main loop into an aligned-K body (EVEN_K=True
  _prefetch_tensors) plus a statically unrolled wind-down for the masked
  tail iterations.
* Runtime-guard the wind-down iters for small num_k_iter so the Final
  iter is the only MFMA that runs when K is short.
* In the main loop, commit the prefetch group before loading scales so
  the compiler schedules buffer_load earlier in the iteration.
* Refresh tuning configs for gfx950.

perf on MI350:

python3 bench_gemm_a8w8_blockscale.py  -gluon

bench_gemm_a8w8_blockscale:
          M       N       K  TFLOPS (Throughput (TFLOPS))
0       1.0  1280.0  8192.0                      0.604139
1      32.0  1280.0  8192.0                     19.064667
2      64.0  1280.0  8192.0                     37.522605
3     128.0  1280.0  8192.0                    100.565860
4     192.0  1280.0  8192.0                     69.512152
5     256.0  1280.0  8192.0                     89.348881
6     320.0  1280.0  8192.0                    115.422745
7     512.0  1280.0  8192.0                    175.689190
8    1024.0  1280.0  8192.0                    345.129363
9    2048.0  1280.0  8192.0                    677.299835
10   4096.0  1280.0  8192.0                    863.537762
11   8192.0  1280.0  8192.0                    887.143030
12  16384.0  1280.0  8192.0                   1164.919752
13   4096.0  4096.0  4096.0                   1271.401835
14   4096.0  4096.0  4160.0                   1076.085957

python3 bench_gemm_a8w8_blockscale.py

bench_gemm_a8w8_blockscale:
          M       N       K  TFLOPS (Throughput (TFLOPS))
0       1.0  1280.0  8192.0                      0.455752
1      32.0  1280.0  8192.0                     13.141420
2      64.0  1280.0  8192.0                     24.324535
3     128.0  1280.0  8192.0                     51.085179
4     192.0  1280.0  8192.0                     85.387665
5     256.0  1280.0  8192.0                    109.271191
6     320.0  1280.0  8192.0                    138.334302
7     512.0  1280.0  8192.0                    218.300780
8    1024.0  1280.0  8192.0                    172.178122
9    2048.0  1280.0  8192.0                    341.678502
10   4096.0  1280.0  8192.0                    670.851040
11   8192.0  1280.0  8192.0                    683.083809
12  16384.0  1280.0  8192.0                    899.010470
13   4096.0  4096.0  4096.0                   1013.235796
14   4096.0  4096.0  4160.0                    862.656740

python3 bench_gemm_a8w8_blockscale.py  -gluon and some non-upstream llvm hack

bench_gemm_a8w8_blockscale:
          M       N       K  TFLOPS (Throughput (TFLOPS))
0       1.0  1280.0  8192.0                      0.554379
1      32.0  1280.0  8192.0                     17.488356
2      64.0  1280.0  8192.0                     34.616803
3     128.0  1280.0  8192.0                     89.166024
4     192.0  1280.0  8192.0                     73.594313
5     256.0  1280.0  8192.0                     97.149177
6     320.0  1280.0  8192.0                    121.220295
7     512.0  1280.0  8192.0                    192.800737
8    1024.0  1280.0  8192.0                    379.143869
9    2048.0  1280.0  8192.0                    742.228581
10   4096.0  1280.0  8192.0                    921.613818
11   8192.0  1280.0  8192.0                    957.073812
12  16384.0  1280.0  8192.0                   1237.227918
13   4096.0  4096.0  4096.0                   1449.159948
14   4096.0  4096.0  4160.0                   1332.912692

* Add GLM GQA FP8 KV paged attention test (#3609)

Co-authored-by: ThomasNing <thomas.ning@amd.com>

* Mhc large m (#3651)

* Add gfx950 large-M hybrid path for mhc_fused_post_pre (M>=1024).

Port EvoKernel big_fuse cache-policy tuning, mhc_post store_nt override,
and hybrid post+gemm+big_fuse dispatch on top of upstream kernels.
Include large-M benchmark script for PR-style perf tables.

* Fix atom docker benchmark to pip install editable aiter.

* Fix mhc_pre_big_fuse_rmsnorm use_nt signature for hybrid compile.

* Add use_nt pybind for mhc_pre_big_fuse_rmsnorm.

* Fix large-M hybrid gemm_out padding to match mhc_pre layout.

* Add gfx950 large-M benchmark results for mhc_large_m branch.

* Update large-M benchmark results and PR comparison tables.

Refresh gfx950 perf data, rename fuse_msnorm to fuse_rmsnorm in output,
and add upstream PR3623 baseline plus PR description markdown.

* Move large-M mhc_post_pre testing to op_tests and drop EvoKernel paths.

Remove docs/EvoKernel benchmark scripts and result artifacts; use
op_tests/test_mhc.py with OOM-safe large-M post_pre sweep instead.

* Refactor gfx950 large-M path into additive kernels only.

Keep PR #3623 kernels unchanged and route M > 1024 to new
mhc_post_pre_large_m / mhc_pre_big_fuse_large_m symbols. Add
op_tests/test_mhc_large_m.py and PR benchmark markdown.

* Fix Black formatting in mhc op tests.

* Remove __main__ guards from mhc op tests.

Match upstream test_mhc.py module-level runner layout and Black formatting.

* Drop PR_mhc_large_m.md and restore op_tests/test_mhc.py to upstream.

Keep large-M coverage in op_tests/test_mhc_large_m.py only.

* Migrate gfx950 large-M cache policy into upstream kernels.

Replace duplicate large-M post/big_fuse kernels with a thin orchestrator
that calls mhc_post and mhc_pre_big_fuse via store_nt/use_nt and
per-stage cache policy resolution, removing ~600 lines of duplication.

* Slim large-M path: drop C++ orchestrator, use mhc_post + mhc_pre in Python.
Remove mhc_post_pre_large_m and big_fuse three-way cache policy; keep post
store_nt RT and large_m_splitk on the force_fused large-M path only.

* Merge gfx950 large-M test into test_mhc.py with --largeM flag.

Remove standalone op_tests/test_mhc_large_m.py; mhc_post_pre summary
adds large_m_us and hip_large_m_err columns on gfx950 when M > 1024.

* Fix Black formatting in test_mhc.py --largeM logging.

* fix(opus): guard gfx942 bf16ws splitk reduce (#3684)

Co-authored-by: yifehuan <yifehuan@local>

* Optimize qk norm rope quant FlyDSL launch path (#3618)

* Optimize qk norm rope quant FlyDSL launch path

* Add flyc.compile + _cf in-memory cache to all FlyDSL kernel launchers

All FlyDSL kernel launchers were calling JitFunction.__call__ directly,
which rebuilds cache keys (sig.bind, DLPack metadata extraction,
dict lookup) on every invocation -- ~170-370 us/call overhead.

Switch to flyc.compile + CompiledFunction._cf caching so that after
the first call, subsequent invocations skip JitFunction overhead and
dispatch directly through the pre-built CallState (~15-25 us/call).

Affected kernels:
  - moe_kernels.py (MoE stage1/stage2)
  - flash_attn_func_gfx1201.py (Flash Attention gfx1201/RDNA4)
  - linear_attention_prefill_kernels.py (chunk_gated_delta_h K5)
  - qk_norm_rope_quant.py (QK norm + RoPE + quant fused)
  - fused_compress_attn.py (fused compress attention)
  - fused_compress_attn_hca.py (HCA compress + norm_rope_scatter)

Co-Authored-By: Claude Opus 4 <noreply@anthropic.com>

* [Bugfix]: _stream_arg

* Unify _cf caching: all kernels use tensor_shim._run_compiled

Consolidate scattered inline _cf / flyc.compile logic into the shared
_run_compiled in tensor_shim.py. Also wire up two previously missed
direct-call sites (moe_gemm_2stage gemm2_exe/reduce_exe, moe_kernels
reduce_exe) and replace fmha_gfx1250's local _run_compiled copy with
an import from tensor_shim.

_cf caching now lives in exactly two places:
  - tensor_shim._run_compiled (general, *args)
  - moe_kernels._run_compiled (moe-specific, tuple + exception cleanup)

Co-Authored-By: Claude Opus 4 <noreply@anthropic.com>

* style: reformat changed files with black (default line-length 88)

Co-Authored-By: Claude Opus 4 <noreply@anthropic.com>

* AOT: precompile moe topk reduction kernel (fix check_aot_cache miss)

The stage2 reduce-mode path runs a separate topk reduction kernel
(launch_moe_reduction) inside the runtime flydsl_moe_stage2 wrapper, but
the AOT script only precompiled the stage2 GEMM directly, so the reduction
kernel was never AOT-covered. It was previously invoked via a direct
launcher call (unmonitored), so check_aot_cache never caught the gap;
routing it through the monitored _run_compiled exposed the JIT-at-runtime
miss.

Fix by extracting the reduction compile+run into a shared
_run_moe_reduction helper in moe_kernels.py, called by both the runtime
stage2 path and the AOT precompile. Sharing the helper guarantees both
derive identical compile-time params (dtype_str / use_mask / num_experts)
and therefore the identical JIT cache key -- no hand-copied logic to drift
out of sync.

AOT covers the single-GPU plain variant (use_mask=False); EP/masked
reduction is a separately-gated multi-GPU path.

Co-Authored-By: Claude Opus 4 <noreply@anthropic.com>

---------

Co-authored-by: xudoyuan <xudoyuan@amd.com>
Co-authored-by: Claude Opus 4 <noreply@anthropic.com>

* fix

* remove

* update kernel name

* update

* flydsl: skip unsupported architectures instead of crashing at import (#3683)

is_flydsl_available() returned True whenever the flydsl package was
installed, even on architectures flydsl ships no kernels for. On such
archs (e.g. gfx1100 / RDNA3) importing aiter.ops.flydsl eagerly registers
GEMM configs and dies with KeyError on flydsl's SMEM_CAPACITY_MAP lookup,
taking down any flash_attn import on that GPU.

Gate availability on the live arch being present in flydsl's
SMEM_CAPACITY_MAP, so flydsl is reported unavailable (and its kernels /
tests are cleanly skipped) on unsupported archs while gfx1151 and the
other supported archs stay enabled.

Changes:
- Membership check uses flydsl's own SMEM_CAPACITY_MAP so the gate
  auto-tracks newly supported archs rather than a hardcoded list.
- Cache the result since arch and package presence are constant per
  process and the check now touches device detection.

* [Triton] Support non-interleaved tensor layout in fused reshape causal conv1d update kernel for Qwen3.5 (#3251)

* Support non-interleaved GQA in fused reshape causal conv1d update

Extends fused_reshape_causal_conv1d_update_single_token (and its
inner Triton kernel) with a gqa_interleaved_layout flag so the
AITER GDN decode fast path can be used for both Qwen3-Next
(interleaved [q,k,v,z] per K-head, [b,a] per K-head group) and
Qwen3.5 (flat [q_all|k_all|v_all|z_all] and [b_all|a_all]) layouts.

The downstream fused_rearrange_sigmoid_gated_delta_rule kernel
already consumes the non-interleaved [q_all|k_all|v_all] conv
output, so no change is needed there.

This lets vllm-project/vllm#42880 drop its gqa_interleaved_layout
guard and route Qwen3.5 through the same fast path that already
serves Qwen3-Next.

Default is gqa_interleaved_layout=True; existing callers are
bit-identical (the branch is a tl.constexpr).

Test: op_tests/triton_tests/test_causal_conv1d_update_single_token.py
parametrized on layout, both pass.

Co-Authored-By: Claude <noreply@anthropic.com>

* use stride check (#3697)

* add cache size

* bench: add fused_rms_mxfp4_quant to model benchmarking tool (#3687)

Register fused RMSNorm + residual add + MXFP4 quant as a kernel in the
triton model benchmarking tool. Reuses bench_rmsnorm.py via a new
--quant mxfp4 / --add-residual mode (no separate bench script), adds a
FusedRmsMxfp4QuantKernelHandler, and shapes for all 7 models in
model_shapes.json. The existing fp4-availability guard auto-gates it to
fp4-capable archs (e.g. gfx950).

* add support for shuffled_kv and blocked version of MLA (#3688)

* [Gluon][GFX950][MLA] Fix mla decode accuracy issue with empty kv split (#3641)

* Fix mla decode empty kv split

* make early return when empty split

* Reduce only loop valid kv seq and guard -inf lse with zero weight

* Enabled stride-aware KV-cache block dim for non-contiguous layouts for fused_qk_norm_rope_cache_pts_quant_shuffle() part 2 (#3640)

* Added strided aware block dim offset calculation to the non shuffle write path and modified unit test for it

* Ruff formatted

* Made non shuffle kv write offet with stride-aware block dim indexing more readable

* Fixed error related to blocksize=0 -> original flat non paged indexing that was missing for CI

* [Gluon] Temporary GPT OSS MoE tuning fix (#3701)

* add

* lazy compile

* Add GLM-4.7-FP8 tuned/untuned BF16 GEMM configs (gfx950) (#3285)

* Adds tuned BF16 GEMM configurations for GLM-4.7-FP8 detected from vLLM server log untuned warnings on gfx950.

* Drop hipBLASLt entries from GLM-4.7 BF16 tuned GEMM config

Per review feedback on #3285: hipBLASLt solution indices
(solidx) are not stable across hipBLASLt builds / docker images.
The same numeric index can resolve to a kernel with mismatched
tile/alignment/split-K constraints in a different image, which in
the worst case causes a GPU hang.

Keep only aiter-native libtypes (asm / skinny / torch) whose solidx
is owned by aiter and stable across builds, matching the convention
used by every other model config in aiter/configs/model_configs/
(dsv3, dsv4, llama70B, glm5, gptoss, kimi, kimik2, ...). Shapes that
previously had only a hipBLASLt winner now fall through to aiter's
runtime defaults in aiter/tuned_gemm.py (heuristic hipBLASLt with
solidx=-1 on gfx942, asm / skinny per the dispatcher otherwise).

Co-authored-by: Cursor <cursoragent@cursor.com>

---------

Co-authored-by: Cursor <cursoragent@cursor.com>

* feat(fmoe): key tuned configs by (gfx, cu_num) to disambiguate archs (#3703)

fmoe tuned configs were keyed on cu_num only, so archs that report the
same CU count (e.g. gfx950 and gfx1250 both 256) collided. Add gfx to the
config key end to end:

- chip_info: gfx_from_cu_num() to backfill legacy cu_num-only rows
  (256->gfx950, 80/304->gfx942; unknown -> runtime arch).
- fused_moe / fused_moe_dp_shared_expert: look configs up by
  (gfx, cu_num, ...) using get_gfx_runtime(); backfill a missing gfx
  column from cu_num so legacy CSVs keep working (no behavior change).
- jit/core: when merging configs, fill a missing gfx column from cu_num
  instead of 0.
- gemm_moe_tune: thread gfx through the tuner key, stamp the real runtime
  arch on tuned rows, and keep gfx as the first CSV column on every write
  (full or incremental).

Co-authored-by: Cursor <cursoragent@cursor.com>

* clean up

* rm hack

---------

Signed-off-by: TennyWang1223 <root@hjbog-srdc-24.amd.com>
Signed-off-by: Aakif Nawaz <aakif.nawaz@amd.com>
Co-authored-by: Matthias Gehre <matthias.gehre@amd.com>
Co-authored-by: la <46212055+junhaha666@users.noreply.github.com>
Co-authored-by: shay-li77 <xiangxli@amd.com>
Co-authored-by: amd-ruitang3 <145657428+amd-ruitang3@users.noreply.github.com>
Co-authored-by: Chengze Fan <fancz2002@gmail.com>
Co-authored-by: yifehuan <Yifeng.Huang@amd.com>
Co-authored-by: yifehuan <yifehuan@local>
Co-authored-by: Xin Huang <Xin.Huang@amd.com>
Co-authored-by: HaonanWang98 <hwang@amd.com>
Co-authored-by: Yutao Xu <xytpai@foxmail.com>
Co-authored-by: yzhou103 <Ying.Zhou2@amd.com>
Co-authored-by: Vinayak Gokhale <Vinayak.Gokhale@amd.com>
Co-authored-by: rocking <ChunYu.Lai@amd.com>
Co-authored-by: Claude Opus 4 (1M context) <noreply@anthropic.com>
Co-authored-by: Jim Guo <jim.guo@amd.com>
Co-authored-by: qichu-yun <qichu@amd.com>
Co-authored-by: azaidy <aliasger.zaidy@amd.com>
Co-authored-by: ahmed-bsod <Muhammad.Ahmed@amd.com>
Co-authored-by: Omar Muhammad <omar.muhammad@amd.com>
Co-authored-by: Alexander Weinrauch <alexander.weinrauch@amd.com>
Co-authored-by: Lukasz Burzawa <lukasz.burzawa@amd.com>
Co-authored-by: Muhammad Ahmed <mm.ahmed2202@gmail.com>
Co-authored-by: honglie <hyi@amd.com>
Co-authored-by: Shao-Chun Lee <Shao-Chun.Lee@amd.com>
Co-authored-by: Mehmet Cagri Kaymak <mehmet.kaymak@amd.com>
Co-authored-by: TennyWang1223 <Tenny.Wang@amd.com>
Co-authored-by: TennyWang1223 <root@hjbog-srdc-24.amd.com>
Co-authored-by: Liyang Ling <leon.ling@amd.com>
Co-authored-by: akii96 <aakif.nawaz@amd.com>
Co-authored-by: ruanjm <jiming.ruan@amd.com>
Co-authored-by: PerryZhang01 <Perry.Zhang@amd.com>
Co-authored-by: perzhang <perzhang@amd.com>
Co-authored-by: inkcherry <mingzhi.liu@amd.com>
Co-authored-by: Kristian Sikiric <kristian.sikiric@amd.com>
Co-authored-by: Juuso Korhonen <40278371+juuso-oskari@users.noreply.github.com>
Co-authored-by: Nidal Danial <81209936+nidal567@users.noreply.github.com>
Co-authored-by: xiaohuguo2023 <149615094+xiaohuguo2023@users.noreply.github.com>
Co-authored-by: Bruno Mazzotti <bruno.mazzotti@amd.com>
Co-authored-by: Satya Nikhil Kodukula <nikhil.kodukula@gmail.com>
Co-authored-by: lijinpei-amd <jinpei.li@amd.com>
Co-authored-by: Thomas Ning <thomasningtrojans@gmail.com>
Co-authored-by: ThomasNing <thomas.ning@amd.com>
Co-authored-by: yinfengLiu <yinfeliu@amd.com>
Co-authored-by: xudoyuan <xudoyuan@amd.com>
Co-authored-by: root <zhimding@amd.com>
Co-authored-by: Zhuo Su <zhuo.su@amd.com>
Co-authored-by: coderfeli <felix.li@amd.com>
Co-authored-by: Jack Hu <jack.hu@amd.com>
Co-authored-by: Olga Miroshnichenko <olga.miroshnichenko@amd.com>
Co-authored-by: Cursor <cursoragent@cursor.com>
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.

4 participants