Skip to content

[ROCm] Add aiter_sage_attn attention backend for gfx1201 (RDNA4) - #1598

Draft
Shan2L wants to merge 2 commits into
ModelTC:mainfrom
Shan2L:feat/rocm-aiter-sage-attn
Draft

Shan2L wants to merge 2 commits into
ModelTC:mainfrom
Shan2L:feat/rocm-aiter-sage-attn

Conversation

@Shan2L

@Shan2L Shan2L commented Oct 11, 2026 •

Copy link
Copy Markdown

Draft: depends on ROCm/aiter#6390, which adds aiter.gfx1201_sage_attention. This PR will be marked ready
after #6390 is merged. MiniMax-H3 on 32 GB GPUs also needs #1595 (allocator config, otherwise OOM at step 1).

This PR adds a SageAttention backend for AMD Radeon RDNA4 (gfx1201: Radeon AI PRO R9700 / R9600D, RX 9070 series) and
a MiniMax-H3 Ref2AV config for 2 x R9700.

Today sage_attn2 and the other Sage backends need the CUDA sageattention packages, so on ROCm MiniMax-H3 runs on
aiter_attn / torch_sdpa, and on gfx1201 on BF16 SDPA.

Changes

  • lightx2v_platform/ops/attn/amd_rocm/aiter_sage_attn.py: platform backend aiter_sage_attn ->
    aiter.gfx1201_sage_attention (INT8 QK^T, FP8 PV, FP32 accumulation; BF16/FP16 Q/K/V, head dim 128, non-causal,
    softmax_scale passed through).
    • Explicit opt-in via attn_type; no default changes.
    • Raises at construction on non-ROCm devices, on GPUs other than gfx1201, or when aiter does not provide the op.
    • Raises on causal, mask, dropout and packed multi-sequence cu_seqlens instead of computing something else.
  • configs/platforms/amd_rocm/minimax_h3_ref2av_r9700_sage.json and
    scripts/platforms/amd_rocm/run_minimax_h3_ref2av_r9700.sh: MiniMax-H3 Ref2AV, TP2, 768x1344, 362 frames,
    29 steps, DiT block offload, AdaLN cache.

CUDA, MI300 / MI350 and the other amd_rocm backends are untouched.

Test

2 x R9700 (gfx1201, 32 GB), ROCm 7.1, PLATFORM=amd_rocm, MiniMax-H3 ref2av TP2, input street_dance
(girl.png, img_0.jpg), seed 42, 768x1344, 362 frames, DiT block offload, AdaLN cache. Only the DiT self-attention
differs between runs.

End to end on main (4fe984c + #1595 + #1596, 3 steps):

attn_type DiT s/step (steps 1/2/3) pipeline total vs torch_sdpa (s/step / total)
torch_sdpa 208.4 / 208.4 / 209.3 913.8 s
aiter_sage_attn 89.0 / 88.9 / 89.0 553.7 s 2.34x / 1.65x
decoded output vs torch_sdpa video PSNR (min frame) SSIM (min frame) audio SNR / cos
aiter_sage_attn 32.28 dB (27.41) 0.959 (0.893) 13.12 dB / 0.975
BF16 SDPA with K/V token order reversed (noise floor, ROCm/aiter#6390, 0.5.0) 35.94 dB (28.62) 0.979 (0.936) 18.91 dB / 0.994

End to end at the stock 29 steps (LightX2V 0.5.0 unmodified, the aiter op installed as the sageattention
package and used through sage_attn2, i.e. the path this backend is bitwise equal to on main; same input, seed and
settings; reproduction script of ROCm/aiter#6390):

attention DiT s/step (mean of steps 2-29) pipeline total vs BF16 SDPA (s/step / total)
BF16 SDPA 211.6 6450 s
gfx1201_sage_attention, HIP kernel (ROCm/aiter#6390) 100.8 3220 s 2.10x / 2.00x
gfx1201_sage_attention, code object (aiter series part 2) 89.1 2877 s 2.38x / 2.24x
sage_attn2 with thu-ml SageAttention (Triton) 347.5 10392 s 0.61x / 0.62x

1450 DiT self-attention calls per rank went through the op, 0 SDPA fallbacks.

The last row is what sage_attn2 can run on gfx1201 today: thu-ml/SageAttention d1a57a5 through
sageattn_qk_int8_pv_fp16_triton (INT8 QK^T, FP16 PV), the kernel sage_attn2 selects on SM89; the CUDA kernels of
the package do not build on ROCm, and its default dispatch for this GPU's capability (12, 0) is a CUDA kernel. It
runs, but is slower than BF16 SDPA.

decoded output pair video PSNR (min frame) SSIM (min frame) audio SNR / cos
code object vs BF16 SDPA 13.59 dB (10.93) 0.578 (0.416) 6.13 dB / 0.877
HIP kernel vs BF16 SDPA 13.04 dB (9.58) 0.551 (0.357) 7.02 dB / 0.899
sage_attn2 with thu-ml SageAttention (Triton) vs BF16 SDPA 14.09 dB (10.20) 0.597 (0.379) 7.13 dB / 0.903
BF16 SDPA K/V reversed vs BF16 SDPA (noise floor) 14.64 dB (10.30) 0.629 (0.391) 9.16 dB / 0.938

Over 29 steps the sampling trajectory is chaotic: reordering the FP32 sums of exact BF16 attention alone already gives a
visibly different video (14.6 dB against the SDPA run), so pixel metrics against one reference run no longer measure attention accuracy; the reference SageAttention
implementation also ends at 14.1 dB against the SDPA run (see the note in
scheduler.py).
Global statistics stay within the spread of the two exact runs:

run luma mean / std sharpness (Laplacian var.) mean frame-to-frame abs diff audio RMS / peak
BF16 SDPA 68.6 / 56.2 40.4 23.24 0.240 / 0.96
BF16 SDPA, K/V reversed 69.0 / 56.7 42.3 23.12 0.233 / 0.91
HIP kernel 69.3 / 57.3 42.6 23.26 0.234 / 0.92
code object 69.3 / 56.9 40.6 23.51 0.238 / 0.91
sage_attn2 with thu-ml SageAttention (Triton) 69.4 / 57.2 41.9 22.92 0.238 / 0.93

29-step frames 0/120/241/361 (columns: BF16 SDPA, BF16 SDPA K/V reversed, HIP kernel, code object):

frames_street_dance_29steps_asm

Before ready for review

Registers the ROCm platform attention backend "aiter_sage_attn", which runs
non-causal BF16/FP16 self-attention through aiter.gfx1201_sage_attention
(INT8 QK^T, FP8 PV, FP32 accumulation). The backend is explicit opt-in via
attn_type and raises at construction on non-ROCm devices, on GPUs other than
gfx1201, and when the installed aiter does not provide the op. Causal
attention, masks, dropout and packed multi-sequence cu_seqlens are rejected.

Signed-off-by: Sylvan Liu <Sylvan.Liu@amd.com>
…h aiter_sage_attn

Signed-off-by: Sylvan Liu <Sylvan.Liu@amd.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.

1 participant