Skip to content

Qwen-Image-2.1 inference on AMD ROCm GPUs (R9700 FP8 / W7900 INT8) - #1549

Merged
helloyongyang merged 22 commits into
ModelTC:mainfrom
zhangnju:rocm_support
Sep 28, 2026
Merged

helloyongyang merged 22 commits into
ModelTC:mainfrom
zhangnju:rocm_support

Conversation

@zhangnju

Copy link
Copy Markdown
Contributor

Summary

Adds native LightX2V inference + optimization for Qwen-Image-2.1 on AMD RDNA GPUs:

  • R9700 (gfx1201 / RDNA4) — fp8-rocm quantization via torch._scaled_mm. ~24 s end-to-end (40 steps), ~17 s (25 steps), single 32 G GPU.
  • W7900 (gfx1100 / RDNA3) — int8-rocm via torch._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 that use_compile is now allowed for qwen_image_21 (default off).

Changes

  • fix(VAE) — im2col GEMM convolution to avoid a nondeterministic MIOpen conv NaN on gfx1201; native conv on gfx1100.
  • fix(sage-attn) — clamp Triton num_stages<=2 on 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.
  • optimize(DiT) — fp8-rocm (torch._scaled_mm) and int8-rocm (torch._int_mm) per-channel-weight / per-token-activation schemes; use_compile via BaseTransformerInfer; optional dit_fuse_qkv_mlp GEMM fusion.
  • Single-GPU configs for R9700 (FP8) and W7900 (INT8), including 25-step practical configs.
  • docs — R9700 / W7900 section in scripts/qwen_image_21/README.md and README_zh.md.

Test plan

  • R9700 fp8 + compile + sage: correct output, ~24 s @40 steps (single GPU)
  • W7900 int8 + compile + sage: correct output, ~46 s @40 steps (single GPU)
  • 25-step configs: no visible quality loss vs 40 steps
  • ROCm logic gated by torch.version.hip / config flags; CUDA default behavior unchanged

zhangnju and others added 22 commits September 22, 2026 01:19
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
helloyongyang merged commit 984f8f3 into ModelTC:main Sep 28, 2026
1 check passed
@zhangnju
zhangnju deleted the rocm_support branch September 29, 2026 00:13
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.

2 participants