Repository navigation
[Triton/Gluon] [gfx950] [dsv4.1-flash] mHC fused kernel - #5824
Conversation
🏷️ CI GuideRuns automatically on every PR:
Extended tests (opt-in via labels):
PR title tags & labels: |
136cd91 to
f746f99
Compare
There was a problem hiding this comment.
Copilot review overview
🟡 Changes recommended
Input-contract bugs, inconsistent hardcoded tuning, and the missing benchmark harness should be addressed before approval.
Get a fresh assessment by requesting another Copilot review.
Review effort: Balanced
Findings: 3
Open (4)
What changed in this PR
Adds a two-launch Triton fusion for delayed mHC seams, reducing intermediate launches in DeepSeek-V4.1 decode and prefill paths.
Changes:
- Adds fused main and reduction kernels.
- Exposes a public wrapper with optional post-mixing.
- Adds numerical and empty-input tests.
| File | Description |
|---|---|
aiter/ops/triton/fusions/mhc_post_pre_delayed.py |
Adds the public wrapper and launch logic. |
aiter/ops/triton/_triton_kernels/fusions/mhc_post_pre_delayed.py |
Implements fused Triton kernels. |
aiter/ops/triton/fusions/__init__.py |
Exports the public operation. |
aiter/ops/triton/_triton_kernels/fusions/__init__.py |
Exports internal kernels. |
op_tests/triton_tests/fusions/test_mhc_post_pre_delayed.py |
Adds correctness coverage. |
💡 Configure MCP servers for context-aware, tailored reviews. Learn more in the docs.
- add bench script - add tuned configs and config picking function - the configs for gfx942 are just copies of gfx950. added them since it mentions in the aiter PR guidelines - address some copilot comments
There was a problem hiding this comment.
Copilot review overview
🟡 Changes recommended
The wrapper can corrupt non-contiguous output buffers and silently ignore incomplete post-mix inputs.
Get a fresh assessment by requesting another Copilot review.
Review effort: Balanced
Findings: 1



Currently in dsv4.1-flash for gfx950,
Prefill (num_tokens >= 1024) uses 5 different kernels for mHC:
mhc_post_kernelmhc_pre_gemm_sqrsummhc_pre_big_fuse_mhc_pre_mix_kernel_hc_head_reduce_store_kerneland decode (num_tokens < 1024) uses 4 kernels for mHC:
mhc_fused_post_pre_gemm_sqrsummhc_pre_big_fuse_mhc_pre_mix_kernel_hc_head_reduce_store_kerneland then a kernel for rmsnorm on the input to the attention/moe block
add_rmsnorm_quant_kernelWe fuse these kernels and replace them with a main kernel (
_mhc_post_pre_delayed_main_kernel) and reduce kernel (_mhc_post_pre_delayed_reduce_kernel).On a per seam basis (mHC work between two sub-layers) the fused kernel has a speedup of 1.16x for prefill and a speedup of 1.47x for decode compared to the unfused kernels (measured from an e2e trace of
ISL 1024/OSL 32/CONC 32/num_prompts 32/max-num-batched-tokens 16384):