Repository navigation
Conversation
…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.
Seas0
requested review from
AgainstEntropy,
BBuf,
HaiShaw,
JustinTong0323,
kevin-mii,
mickqian,
ping1jing2,
sogalin,
wisclmy0611,
yichiche,
yingluosanqian and
zijiexia
as code owners
September 23, 2026 15:29
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.
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:
runtime/models/vaes/qwen_image21_vae_cuda_opt.py, registered inplatforms/cuda.py: aVaeFastPathGate-controlled fast path for encoder and decoder. With--quality extra-highorhighevery convolution takes channels_last input and lets cuDNN apply its own zero padding (noF.padcopy, no transposes), everyRMS_norm -> SiLUchain 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.Accuracy Tests
test_qwen_image21_cuda.py::test_vae_fast_path_keeps_names_and_lossless_output: state-dict names unchanged, gate-off decode and encodetorch.equalto 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_silupasses including the new 1152-channel case.--quality extra-highagainst 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-highvs lossless: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