Repository navigation
Conversation
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>
This was referenced Oct 11, 2026
This file contains hidden or bidirectional Unicode text that may be interpreted or compiled differently than what appears below. To review, open the file in an editor that reveals hidden Unicode characters.
Learn more about bidirectional Unicode characters
Sign up for free
to join this conversation on GitHub.
Already have an account?
Sign in to comment
Add this suggestion to a batch that can be applied as a single commit.This suggestion is invalid because no changes were made to the code.Suggestions cannot be applied while the pull request is closed.Suggestions cannot be applied while viewing a subset of changes.Only one suggestion per line can be applied in a batch.Add this suggestion to a batch that can be applied as a single commit.Applying suggestions on deleted lines is not supported.You must change the existing code in this line in order to create a valid suggestion.Outdated suggestions cannot be applied.This suggestion has been applied or marked resolved.Suggestions cannot be applied from pending reviews.Suggestions cannot be applied on multi-line comments.Suggestions cannot be applied while the pull request is queued to merge.Suggestion cannot be applied right now. Please check back later.
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_attn2and the other Sage backends need the CUDAsageattentionpackages, so on ROCm MiniMax-H3 runs onaiter_attn/torch_sdpa, and on gfx1201 on BF16 SDPA.Changes
lightx2v_platform/ops/attn/amd_rocm/aiter_sage_attn.py: platform backendaiter_sage_attn->aiter.gfx1201_sage_attention(INT8 QK^T, FP8 PV, FP32 accumulation; BF16/FP16 Q/K/V, head dim 128, non-causal,softmax_scalepassed through).attn_type; no default changes.cu_seqlensinstead of computing something else.configs/platforms/amd_rocm/minimax_h3_ref2av_r9700_sage.jsonandscripts/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_rocmbackends are untouched.Test
2 x R9700 (gfx1201, 32 GB), ROCm 7.1,
PLATFORM=amd_rocm, MiniMax-H3 ref2av TP2, inputstreet_dance(
girl.png,img_0.jpg), seed 42, 768x1344, 362 frames, DiT block offload, AdaLN cache. Only the DiT self-attentiondiffers between runs.
aiter.gfx1201_sage_attentioncall (3D/4D inputs,softmax_scale): bitwise identical.main4fe984c + Set PYTORCH_CUDA_ALLOC_CONF before platform init #1595 + [ROCm] Opt-in deterministic FP32 GEMM/conv via rocBLAS #1596,aiter_sage_attnvssage_attn2with the same aiter op installed asthe
sageattentionpackage: video and audio bitwise identical (ROCM_DETERMINISTIC_FP32_BLAS=1, [ROCm] Opt-in deterministic FP32 GEMM/conv via rocBLAS #1596).ruff check/ruff format --check(v0.11.0) clean.End to end on
main(4fe984c + #1595 + #1596, 3 steps):attn_typetorch_sdpa(s/step / total)torch_sdpaaiter_sage_attntorch_sdpaaiter_sage_attnEnd to end at the stock 29 steps (LightX2V 0.5.0 unmodified, the aiter op installed as the
sageattentionpackage and used through
sage_attn2, i.e. the path this backend is bitwise equal to onmain; same input, seed andsettings; reproduction script of ROCm/aiter#6390):
gfx1201_sage_attention, HIP kernel (ROCm/aiter#6390)gfx1201_sage_attention, code object (aiter series part 2)sage_attn2with thu-ml SageAttention (Triton)1450 DiT self-attention calls per rank went through the op, 0 SDPA fallbacks.
The last row is what
sage_attn2can run on gfx1201 today: thu-ml/SageAttentiond1a57a5throughsageattn_qk_int8_pv_fp16_triton(INT8 QK^T, FP16 PV), the kernelsage_attn2selects on SM89; the CUDA kernels ofthe 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.
sage_attn2with thu-ml SageAttention (Triton) vs BF16 SDPAOver 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:
sage_attn2with thu-ml SageAttention (Triton)29-step frames 0/120/241/361 (columns: BF16 SDPA, BF16 SDPA K/V reversed, HIP kernel, code object):
Before ready for review
scripts/platforms/amd_rocm/run_minimax_h3_ref2av_r9700.shas shipped (29 steps, with the AdaLN cachestep) on the merged aiter and post the numbers here