Qwen-Image-2.1 inference on AMD ROCm GPUs (R9700 FP8 / W7900 INT8) - #1549
Merged
Merged
Conversation
MIOpen convolution kernels return NaN nondeterministically for the Qwen-Image-2.1 VAE on gfx1201 (RDNA4). Route every VAE convolution (ImageConv, the Resample upsampler, and AttentionBlock) through im2col + GEMM (hipBLASLt), which is numerically reliable there. Add opt-in vae_conv_im2col plus vae_device / vae_dtype so the VAE can also run on CPU or in a chosen dtype as a fallback.
… gfx1201
SageAttention's INT8 WMMA Triton kernels fail to compile on ROCm/gfx1201
at num_stages>=3 ("operation destroyed but still has uses" in the AMD
pipeliner pass). Clamp num_stages<=2 on both SageAttention's own kernels
and every torch.compile / Inductor Triton compilation, and run Inductor
single-threaded so the clamp reaches its worker (only on the SageAttention
path). With this, dense SageAttention2 runs and is faster than torch SDPA
flash on this hardware.
…ile, cross-GPU text encoder, QKV/MLP fusion - fp8-rocm quant scheme: per-channel weight + per-token activation FP8 via torch._scaled_mm, auto-quantized from the released BF16 weights (no prequantized checkpoint; releases the BF16 source on load to halve the load-time peak). CUTLASS/sgl FP8 kernels are CUDA-only. - use_compile for qwen_image_21 via BaseTransformerInfer (infer_block / run_block); the block graph passes only tensors, keeping modulation and KV-cache lookups outside it. - Run the text encoder on a separate GPU (text_encoder_device_index): ROCm empty_cache does not return the onloaded ~15G to the driver, which would otherwise starve the VAE decode on the main GPU. - Optional dit_fuse_qkv_mlp: fuse q/k/v and gate/up into single FP8 GEMMs (numerically identical; one shared activation quantization).
Torch-op inference path for RDNA4 (torch_sdpa, torch_real_rope, im2col VAE, text encoder on cuda:1): - qwen_image_21_r9700.json : bf16 baseline - ..._fp8_compile.json : fp8 + torch.compile - ..._fp8_compile_sage.json : fp8 + compile + SageAttention2 (fastest, 40 steps) - ..._fp8_compile_fused.json : + QKV/MLP fusion (opt-in) - ..._fp8_sage_25steps.json : 25-step practical config (~1.6x fewer steps, no quality loss)
…00 / W7900) RDNA3 has no FP8 tensor path, so add an INT8 scheme mirroring fp8-rocm: - int8-rocm: per-channel weight + per-token activation INT8 via torch._int_mm, auto-quantized from the released BF16 weights (no prequantized checkpoint). Falls back to a BF16 GEMM for the small prefill segments where torch._int_mm requires M > 16. - dit_fuse_qkv_mlp now also fuses the int8 q/k/v and gate/up GEMMs. - Allow int8-rocm in the qwen_image_21 runner and base-model quant whitelists. On gfx1100 native MIOpen VAE conv is fine (no im2col needed) and SageAttention2 INT8 is ~1.75x over torch SDPA flash, so the practical W7900 stack is int8-rocm + torch.compile + SageAttention2.
Native-conv VAE + CPU-offload text encoder (no gfx1201 workarounds): - qwen_image_21_w7900_int8.json : int8 baseline - ..._int8_compile.json : int8 + torch.compile - ..._int8_compile_sage.json : int8 + compile + SageAttention2 (fastest, 40 steps) - ..._int8_sage_25steps.json : 25-step practical config (~29s end-to-end)
Drop text_encoder_device_index and use text_encoder_cpu_offload instead. The im2col GEMM VAE decode reuses PyTorch's allocator pool (where the offloaded text-encoder memory is cached), so the VAE no longer starves for MIOpen workspace on the main GPU. A single 32G R9700 fits the whole pipeline (verified: bf16 and fp8, same speed as the previous two-GPU setup).
Add a section (in both README.md and README_zh.md) covering the AMD RDNA runs: deps, single-GPU run command, per-GPU configs (fp8-rocm on R9700 / int8-rocm on W7900), ROCm-specific notes (BF16-on-load quantization, im2col VAE on gfx1201, SageAttention num_stages workaround), and measured end-to-end latency at 40/25 steps.
…_platform Address review: keep platform-specific code under the platform layer. - MM: fp8-rocm/int8-rocm GEMM (torch._scaled_mm / torch._int_mm) moved to lightx2v_platform/ops/mm/amd_rocm and registered via PLATFORM_MM_WEIGHT_REGISTER; removed from the shared lightx2v/common/ops/mm/mm_weight.py. - Attn: SageAttention/Triton ROCm num_stages workaround moved to lightx2v_platform/ops/attn/amd_rocm/sage_patch.py; removed from the shared sage_attn.py. - Configs moved to configs/platforms/amd_rocm; run under PLATFORM=amd_rocm. - Make aiter optional in AmdRocmDevice so the torch-native FP8/INT8 path needs no aiter build (aiter is not required on R9700/W7900). Validated end-to-end (1024x1024, 40 steps, seed 42): R9700 fp8 ~24.4s, W7900 int8 ~44.9s; both produce correct images.
…cope, add tests - VAE: drop the im2col/GemmConv2d workaround and vae_device/vae_dtype; the AMD platform already disables cuDNN/MIOpen, so gfx1201 runs native conv without the nondeterministic NaN (verified on R9700). vae.py returns to plain nn.Conv2d. - Remove QKV / gate-up fusion (measured ~0%, compute-bound) and its AMD-specific fused MM classes + config; the model no longer imports AMD backend classes and the infer path no longer branches on hasattr(block, "qkv"). - SageAttention Triton num_stages workaround now installs only when sage_attn2 is selected (from SageAttn2Weight.__init__) and only on validated archs (gfx1201/gfx1100), instead of unconditionally for every ROCm workload. - fp8/int8 MM: fail-fast if torch._scaled_mm/_int_mm is missing or arch is wrong; int8 GEMM inputs made contiguous. - Configs/README: rename *_sage_25steps -> *_compile_sage_25steps, relabel INT8 eager baseline, drop the experimental fused config, pin tested transformers version, soften the 25-step quality claim, add benchmark methodology. - Add lightx2v_platform/test/test_amd_rocm_quant.py (fp8/int8 vs bf16 across M=1/16/17/31/64, incl. the int8 fallback boundary; registration without aiter). Validated on R9700 (fp8) and W7900 (int8): quant test ALL PASS, t2i produces correct images (native conv, no fusion), 1024x1024.
…rm layer + add run scripts
- Move the SageAttention/Triton num_stages workaround entirely into the amd_rocm
platform ops (applied when amd_rocm attention loads, arch-gated to gfx1201 /
gfx1100). Removes the backend (torch.version.hip) branch from the shared
SageAttn2Weight.__init__, keeping shared inference code backend-agnostic.
- Add scripts/platforms/amd_rocm/qwen_image_21_{r9700,w7900}_t2i.sh (export
PLATFORM=amd_rocm + source scripts/base/base.sh), matching the other platforms;
point the README at them instead of an inline PLATFORM= command.
Re-validated on R9700: patch applies at platform load, fp8+compile+sage 25 steps
~15.4s, output unchanged.
… slim configs, real pytest - SageAttention Triton workaround now fires only when a sage_attn2 backend is actually constructed: shared SageAttn2Weight calls a uniform platform hook (get_platform_device().on_sage_attn2_init), which the amd_rocm device implements (arch-gated). No import-time side effect on non-sage ROCm runs, and no backend branch in shared inference code. - Revert the multi-GPU text_encoder_device_index change in the Qwen runner (unused, backend-specific); drop the now-dead `import contextlib`. - Revert the incidental linear_keys / mlp refactors so the Qwen transformer files differ from base only by the (necessary) torch.compile support. - Reduce configs to 3 user-facing ones: qwen_image_21_r9700_fp8 / _w7900_int8 / _bf16; point the two run scripts + README at them. - mm: _gcn_arch() uses current_device(); fp8 error no longer prescribes INT8; sage patch logs skipped modules at debug. - Rewrite the quant test as real pytest (collectable, skips off-ROCm), exercising the production MM_WEIGHT_REGISTER with a fixed seed. Validated: pytest passes on R9700 (3) and W7900 (2, fp8 skipped); t2i correct on both (R9700 ~24s, W7900 ~47s) via the run scripts.
… (10 args -> 7) The four modulation tensors were split into separate params for the (now removed) fusion path; the base code already passed them as one tuple. Regroup them so the compiled block takes modulation/rotary/positions/k_cache/v_cache. Re-validated on R9700 (fp8+compile+sage, ~24s).
… tighten helpers) Functions are private and only run after apply_rocm_sage_patches() has already checked the arch, so the per-function hip guards were redundant; use min() / startswith(tuple) / a _MAX_STAGES constant and trim docstrings. Behavior unchanged — re-validated on R9700 (num_stages clamp on 15 kernels, ~24s).
…v var LIGHTX2V_ROCM_TRITON_MAX_STAGES (default 2) is the clamp ceiling; <=0 disables the whole workaround with no code change (e.g. once Triton fixes the pipeliner bug). Documented in the README. Env var (not JSON config) because the patch is a process-level, one-time monkeypatch applied before any per-request config. Re-validated on R9700 (default path still clamps 15 kernels, ~24s).
helloyongyang
approved these changes
Sep 28, 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.
Summary
Adds native LightX2V inference + optimization for Qwen-Image-2.1 on AMD RDNA GPUs:
fp8-rocmquantization viatorch._scaled_mm. ~24 s end-to-end (40 steps), ~17 s (25 steps), single 32 G GPU.int8-rocmviatorch._int_mm(RDNA3 has no FP8). ~46 s (40 steps), ~29 s (25 steps), single 48 G GPU.Measured at 1024x1024, seed 42, CFG off, steady state after warmup. Weights are quantized from the released BF16 checkpoint on load (no converter, no
dit_quantized_ckpt).All ROCm-specific behavior is gated by
torch.version.hip, config flags, or quant-scheme selection, so CUDA paths are unchanged by default. The one cross-platform change is thatuse_compileis now allowed forqwen_image_21(default off).Changes
num_stages<=2on ROCm (SageAttention kernels + Inductor) to dodge the AMD pipeliner crash; run Inductor single-threaded on the SageAttention path. Working SageAttention2 (INT8) is ~1.13x over flash on gfx1201 and ~1.75x on gfx1100.fp8-rocm(torch._scaled_mm) andint8-rocm(torch._int_mm) per-channel-weight / per-token-activation schemes;use_compileviaBaseTransformerInfer; optionaldit_fuse_qkv_mlpGEMM fusion.scripts/qwen_image_21/README.mdandREADME_zh.md.Test plan
torch.version.hip/ config flags; CUDA default behavior unchanged