Skip to content

[diffusion] perf: quality-gated channels_last fast path for the Qwen-Image 2.1 VAE - #40934

Closed
Seas0 wants to merge 3 commits into
sgl-project:mainfrom
Seas0:qi21/vae-fast-path
Closed

Seas0 wants to merge 3 commits into
sgl-project:mainfrom
Seas0:qi21/vae-fast-path

Conversation

@Seas0

@Seas0 Seas0 commented Sep 23, 2026 •

Copy link
Copy Markdown
Contributor

Motivation

The Qwen-Image 2.1 VAE is its own class and had none of the Wan-family fast paths. Its decode profile at 1024x1024 was 52% convolutions, 13% NCHW/NHWC transposes that cuDNN inserts around every conv because activations stay NCHW, 8% separate conv-bias adds, and the rest FP32 copies for the channel RMSNorm, explicit padding copies and SiLU. Decode becomes a 3% share of a request once Cache-DiT is on, and more on consumer GPUs.

Modifications

Three commits:

  1. runtime/models/vaes/qwen_image21_vae_cuda_opt.py, registered in platforms/cuda.py: a VaeFastPathGate-controlled fast path for encoder and decoder. With --quality extra-high or high every convolution takes channels_last input and lets cuDNN apply its own zero padding (no F.pad copy, no transposes), every RMS_norm -> SiLU chain runs the fused channels_last_3d Triton kernel, and the nearest-exact upsample runs the NHWC gather. The lossless default runs the original modules bit for bit; class swaps keep every parameter under its checkpoint name. The 2.1 encoding stage toggles the gate around VAE encode the way the decode stage already does.
  2. The fused Wan RMSNorm+SiLU kernel accepts up to 2048 channels so the decoder's 1152-channel levels fuse (numerics test extended with a 1152-channel case). Note: [Diffusion] Fuse the Wan VAE conv bias epilogue and tile the RMSNorm+SiLU kernel #38650 re-tiles this kernel and keeps the 1024 cap; the two will conflict textually and this one-line change is easy to carry over.
  3. Cookbook note and the benchmark skill's fast-path inventory.

Accuracy Tests

  • test_qwen_image21_cuda.py::test_vae_fast_path_keeps_names_and_lossless_output: state-dict names unchanged, gate-off decode and encode torch.equal to the uninstalled model, gate-on outputs close, the gate resets after the block, and the fused norm sees channels_last_3d input. test_norm.py -k wan_rmsnorm_silu passes including the new 1152-channel case.
  • In the pipeline, --quality extra-high against the lossless image of the same seed: A800 edit SSIM 0.9988 / PSNR 43.1 dB, RTX 4090 edit SSIM 0.9989 / 43.2 dB; on text-to-image the VAE path adds no measurable drift at four digits.

Speed Tests and Profiling

1024x1024, synchronized stage timing, --quality extra-high vs lossless:

A800 RTX 4090
decode stage (ms) 181 -> 128 246 -> 172
edit encoding stage (ms) 259 -> 240 noisy on that host (encoder streamed)
peak allocated, t2i (GB) 34.3 -> 33.3 18.8 -> 17.8

Decode after: convs 67%, separate conv-bias adds 10%, fused norm 5%. The remaining opportunity is the conv-bias epilogue that #38650 adds for the Wan VAE.

Checklist


CI States

Latest PR Test (Base): ❌ Run #35881994679
Latest PR Test (Extra): ❌ Run #35881994493
Latest PR Test (AMD ROCm 10): ❌ Run #35881994508

…Image 2.1 VAE

Install a VaeFastPathGate-controlled fast path on the 2.1 VAE encoder and
decoder at load time. With quality=extra-high or high every convolution
takes channels_last input and lets cuDNN apply its own zero padding (no
explicit F.pad copy, no NCHW/NHWC transposes around each conv), every
RMS_norm -> SiLU chain runs the fused channels_last_3d Triton kernel, and
the 2D nearest-exact upsample runs the NHWC gather. The lossless default
runs the original modules bit for bit; class swaps keep every parameter
under its checkpoint name. The 2.1 encoding stage now toggles the gate
around VAE encode the way the decode stage already does.

A800, 1024x1024, in the pipeline: decode 181 -> 135 ms, edit encoding
259 -> 240 ms, peak memory -1 GB; edit output SSIM 0.9988 / PSNR 43.1 dB
against the lossless image.
…annels

The kernel normalizes one pixel's channel row per program with a
next-power-of-two block, so widths above 1024 only need the 8-warp block it
already selects. Raise the cap so the Qwen-Image 2.1 decoder's 1152-channel
levels use the fused path instead of the eager norm plus SiLU (decode
130 -> 124 ms on A800), and add the 1152-channel case to the numerics test.
Note in the cookbook that --quality extra-high or high enables the VAE
channels_last fast path for encode and decode with the measured A800
numbers, and add the path to the benchmark skill's fast-path inventory
with the remaining conv-bias opportunity.
Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Labels

diffusion SGLang Diffusion documentation Improvements or additions to documentation jit-kernel

Projects

None yet

Development

Successfully merging this pull request may close these issues.

1 participant