[triton-mha] hint head-stride div-by-8 for vectorized global load - #3424
Conversation
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.
🏷️ CI GuideRuns automatically on every PR:
Extended tests (opt-in via labels):
|
There was a problem hiding this comment.
Pull request overview
Improves performance of the default Triton MHA forward kernel by enabling vectorized global loads when Q/K/V head-axis strides are 8-element aligned (not just 16), which helps common packed [seq, heads, dim] layouts such as varlen prefill with head_dim like 72.
Changes:
- Adds a
HEAD_STRIDE_ALIGNED_8tl.constexprto_attn_fwdand appliestl.multiple_of(..., 8)to the head-axis offset to improve alignment inference. - Updates the Python wrapper to compute and pass
HEAD_STRIDE_ALIGNED_8based on runtime Q/K/V head-axis strides (stride_*h % 8 == 0).
Reviewed changes
Copilot reviewed 2 out of 2 changed files in this pull request and generated 1 comment.
| File | Description |
|---|---|
aiter/ops/triton/attention/mha.py |
Computes and passes HEAD_STRIDE_ALIGNED_8 meta-parameter to the Triton kernel based on runtime strides. |
aiter/ops/triton/_triton_kernels/attention/mha.py |
Adds the new constexpr and applies tl.multiple_of hints on head-axis offsets to encourage vectorized loads. |
💡 Add Copilot custom instructions for smarter, more guided reviews. Learn how to get started.
brunomazzottiamd
left a comment
There was a problem hiding this comment.
Hello @mgehre-amd. The code is good, I have no objections to it.
However, I'm a bit concerned about the restricted benchmark. How can we be sure that no performance regressions are introduced for other attention flavors and GPU archs? Am I missing anything?
Can you please expand the benchmarking cases? I can think of causal / non-causal, bshd / varlen layouts, PE variant, other GPU archs (gfx942, gfx950). Do you have access to CDNA GPUs?
Please take a look at op_tests/op_benchmarks/triton/model_benchmarking_tool/bench_attn_models.py script. I can't remember if it can cover all cases, but it's worth checking it, probably it will help to extend the benchmark.
|
@micmelesse, can you please help reviewing this PR? The author said it "Mirrors the equivalent hint added to the flash_attn_triton_amd (dao_ai) prefill kernel". Your knowledge about DAO AI / FlashAttention implementations will help a lot here. Thanks in advance! |
For clarity, "Mirrors the equivalent hint" refers to my other PR #3419 |
@brunomazzottiamd, thanks for your feedback! I was wondering about the best way to validate this. Which validation has been run on other PRs targeting this kernel? I will checkout op_tests/op_benchmarks/triton/model_benchmarking_tool/bench_attn_models.py and let me see if I have access to some gfx942/gfx950 internally. |
@mgehre-amd, the only automated way I can think right now is Unfortunately we don't have an automated performance CI and we have been trusting performance reports from PR authors until now. I can help you collect performance data on |
|
Thanks for the review! Here's an expanded benchmark covering causal/non-causal, bshd/varlen layouts, and multiple Regression safety argument
Extended benchmark — gfx1151 (Strix Halo / RDNA3.5)Toolchain: torch 2.11.0+rocm7.14.0a20260529, triton 3.7.0 (built from
Summary for gfx1151:
Extended benchmark — gfx950 (MI350X)Toolchain: torch 2.9.1+rocm7.14.0a20260609, triton 3.5.1+rocm7.14.0a20260609.
Summary for gfx950 (5-rep medians):
Extended benchmark — gfx942 (MI300X)Toolchain: torch 2.9.1+rocm7.14.0a20260609, triton 3.5.1+rocm7.14.0a20260609.
Summary for gfx942: |
brunomazzottiamd
left a comment
There was a problem hiding this comment.
Given the evidence presented in #3424 (comment), I withdraw my cautionary objection to this PR. Let's see if anyone else has something to add...
@mgehre-amd, can you please do a quick bench of PE case? It's used in DeepSeek-V{2,3} prefill workload.
BSHD layout:
python op_tests/op_benchmarks/triton/bench_mha.py \
-impl default -metric time \
-causal true --dtype bf16 \
-b 1 -hq 128 -hk 128 -d 192 -dv 128 -sq 4096 -sk 4096 \
-fn fwdTHD (varlen) layout:
python op_tests/op_benchmarks/triton/bench_mha.py \
-impl default -metric time \
-causal true --dtype bf16 \
-b 1 -hq 128 -hk 128 -d 192 -dv 128 -sq 4096 -sk 4096 \
-fn fwd_varlen -equal_seqlens|
Shape: B=1, HQ=HK=128, d=192, dv=128, sq=sk=4096, causal=True, bf16.
|
brunomazzottiamd
left a comment
There was a problem hiding this comment.
LGTM! No performance regressions at all.
* 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>
For the packed
[seq, heads, dim]layout used by varlen prefill, the head-axis stride equalshead_dim. Whenhead_dimis a multiple of 8 but not 16 (e.g. 72), Triton's integer-arg auto-specialization does not attachtt.divisibility = 8tostride_*h(its threshold is 16), so AxisInfo treats the K/V global load as 2-byte aligned and Coalesce emits scalarbuffer_load_u16instead of vectorizedbuffer_load_b128.Add a
HEAD_STRIDE_ALIGNED_8constexpr to_attn_fwdand applytl.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 == 0against the actual runtime strides (not againsthead_dim), so the hint stays sound for non-contiguous Q/K/V views wherestride_*h != head_dim. The constexpr defaults toFalse, so external callers of_attn_fwdare 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 #3423 this closes the gap with the ck MHA on Strix Halo.Benchmark — gfx1151,
bench_mha -impl defaultQwen3-Omni ViT prefill shape: B=1, S=3200, H=16, head_dim=72, fp16, varlen, non-causal. Toolchain: torch 2.11.0+rocm7.14.0a20260529, triton 3.7.0 (built from
triton-lang/tritonmain,d92727c2).bench_mhatimeReproducer
python op_tests/op_benchmarks/triton/bench_mha.py \ -fn fwd_varlen -equal_seqlens \ -b 1 -hq 16 -hk 16 -sq 3200 -sk 3200 -d 72 \ --dtype fp16 -causal False \ -impl default -metric time