Repository navigation
[Diffusion] Keep the Wan VAE decoder channels_last and add a Triton NHWC nearest upsample - #38182
Merged
Merged
Conversation
…HWC nearest upsample Follow-up to sgl-project#38020. The degenerate-stride layout issue found on the Qwen-Image VAE also affects the Wan 2.1 video VAE, where the decoder is a much larger share of short-step video pipelines, and it comes with a second problem: once the decoder does run channels_last, aten's upsample_nearest2d_nhwc kernel is several times slower than its NCHW sibling on the same bytes. Lossless (bit-exact, unconditional): - New kernel layout/nearest_upsample_nhwc_triton.py: integer-factor nearest upsample of a dense channels_last [N, C, H, W] tensor as a pure gather. For integer factors nearest and nearest-exact both map output index i to input index i // f, so the result is bitwise identical to nn.Upsample in either mode for every dtype. Registered as diffusion.nearest_upsample_nhwc with a can_use_ predicate. 0.061 ms vs 0.458 ms (aten NHWC) and 0.190 ms (aten NCHW) for [1, 192, 240, 416] bf16 on H200. - GatedChannelsLastUpsample runs it whenever the upsample input already has canonical NHWC strides (aten would run its NHWC kernel and return the same dense channels_last tensor), on the lossless path too. Quality-gated (extra-high / high): - resample_forward's upsample3d frame interleave (reshape / stack / reshape, which materialises NCDHW) writes the same values straight into a channels_last_3d buffer with one copy via _interleave_time_pairs. - The Wan installer now wraps WanUpsample with GatedChannelsLastUpsample and hands the gate to WanResample and WanAttentionBlock (residual operand order), so all 29 RMSNorm+SiLU sites hit the fused kernel instead of 25. The attention operand order is commutative but changes the reduction order of the eager norm that follows under fp32 autocast, which is how Wan pipelines decode (fp32 VAE weights, bf16 autocast); measured end-to-end as non-bit-exact, hence gated. Wan 2.1 VAE, 81 frames 480x832 on H200, torch.profiler GPU kernel time: bf16 weights: main lossless 1342 ms / high 1178 ms -> 1324 / 1112 ms. fp32 weights + bf16 autocast (the T2V pipeline setting): main 1633 / 1311 ms -> 1593 / 1215 ms. Lossless mp4 from TurboWan2.1-T2V-1.3B md5-identical to main; high vs main-high PSNR 36.3 dB mean / 35.0 dB min over 81 frames. Peak memory unchanged (4.577 GB). Qwen-Image high: 44.6 -> 42.9 ms.
Dayuxiaoshui
requested review from
AgainstEntropy,
BBuf,
HaiShaw,
mickqian,
ping1jing2,
yichiche and
yingluosanqian
as code owners
September 6, 2026 04:27
BBuf
reviewed
Sep 7, 2026
BBuf
reviewed
Sep 7, 2026
BBuf
reviewed
Sep 7, 2026
…contract Review follow-ups: - can_use_nearest_upsample_nhwc never raises: _integer_scale rejects non-numeric entries, NaN and +/-Inf (and bools) instead of throwing. - The predicate now requires C > 1 and canonical NHWC strides on every dim. That is exactly the condition under which aten's suggest_memory_format picks channels_last, so the Triton result is layout-identical as well as value-identical; with C == 1 the tensor is also NCHW-contiguous and aten returns e.g. strides (4, 4, 2, 1) for a [1, 1, 2, 2] output. - nearest_upsample_nhwc validates before launching and raises ValueError for requires_grad inputs (a direct call must not return a detached tensor), unsupported dtype/device/rank, and non-canonical layouts. - GatedChannelsLastUpsample relies on the predicate alone for the lossless dispatch instead of duplicating the stride check. - Tests assert out.stride() == ref.stride(), cover the malformed scale factors, the C == 1 case and the direct-call error paths.
BBuf
approved these changes
Sep 8, 2026
Collaborator
|
/tag-and-rerun-ci |
5 tasks done
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
Follow-up to #38020. The degenerate-stride layout issue found on the Qwen-Image VAE also affects the Wan 2.1 video VAE, where the decoder is a much larger share of a request: on TurboWan2.1-T2V-1.3B (4 steps, 81 frames, 480x832) a warm request spends ~0.93 s denoising and ~1.56 s in
DecodingStage, i.e. the VAE is ~60% of the request.Two things were going on in the Wan decoder at
quality=high:resample_forward's upsample3d frame interleave (reshape / stack / reshape) materialises NCDHW, and the attention residual's operand order drops channels_last_3d too, so every up block started with an NCDHW tensor: 4 of the 29RMSNorm+SiLUsites fell back to eager per chunk, the residual adds ran strided, and cuDNN converted the resample conv2d input to NHWC and back.upsample_nearest2d_nhwckernel is several times slower than its NCHW sibling on the same bytes (0.458 ms vs 0.190 ms for[1, 192, 240, 416]bf16 on H200); the decoder hits it once per up block per chunk.Modifications
New kernel, lossless (bit-exact, unconditional)
kernels/ops/diffusion/layout/nearest_upsample_nhwc_triton.py: integer-factor nearest upsample of a dense channels_last[N, C, H, W]tensor as a pure gather. For an integer factorf, writingi = k*f + rwith0 <= r <= f-1,nearestreadsfloor(i/f) = kandnearest-exactreadsfloor((i+0.5)/f) = ksince(r+0.5)/f < 1, so both readi // fand the result is bitwise identical tonn.Upsamplein either mode. Walks the output in NHWC memory order so loads and stores are contiguous alongC. 0.061 ms for the shape above. Registered asdiffusion.nearest_upsample_nhwcwithcan_use_nearest_upsample_nhwc; raises on unsupported input.GatedChannelsLastUpsampleruns it whenever the upsample input already has canonical NHWC strides. In that case aten itself would run its NHWC kernel and return the same dense channels_last tensor, so the replacement is layout- and value-identical and is enabled on the lossless path too.Quality-gated (
extra-high/high)wanvae.py:_interleave_time_pairswrites the interleaved frames straight into achannels_last_3dbuffer with onecopy_(both views are valid: the channel dim has stride 1 so[B, 2C, ...] -> [B, 2, C, ...]is a view; the output's2Tdim splits into(T, 2)), replacing the NCDHW-materialisingstack. Gated because the NHWC conv2d it enables need not pick the same cuDNN algorithm.WanUpsamplewithGatedChannelsLastUpsample(stride canonicalisation, from [Diffusion] Port the Wan VAE decoder fast paths to the Qwen-Image VAE #38020) and hands the gate toWanResampleandWanAttentionBlock(residual operand order). All 29 norm sites now hit the fused kernel. The operand order is commutative but changes the reduction order of the eager norm that follows under fp32 autocast, which is how Wan pipelines decode (fp32 VAE weights, bf16 autocast); measured end-to-end as non-bit-exact, hence gated.Tests / docs
test_layout.py:nearest_upsample_nhwcvsF.interpolatefor bf16/fp16/fp32, factors 2 and (3, 2),nearestandnearest-exact, N = 1/2/4; predicate rejections.test_model_fast_paths.py: interleave vsstack(values and layout), upsample wrapper dispatch on canonical / degenerate / NCHW inputs with gate on and off, installer wiring.existing-fast-paths.mdupdated.Accuracy Tests
1x H200, Wan 2.1 VAE (
Wan-AI/Wan2.1-I2V-14B-480P-DiffusersVAE weights), 81 frames 480x832.Lossless path
torch.equalvsmainoutput: True.vae_precision=fp32, vae_decode_precision=bf16):torch.equalvsmain: True.sglang generateon TurboWan2.1-T2V-1.3B, 6 prompts, seed 1, 4 steps: all 6 mp4s md5-identical tomain.Quality path
main-high vs this-PR-high (DiT path identical and confirmed deterministic across runs), 81 frames: PSNR mean 36.3 dB, min 35.0 dB. The extra difference vsmain-high comes from the 4 additional norm sites per chunk now reaching the pre-existingwan_rmsnorm_silukernel; the layout-only part was measured at 89.6 dB on Qwen-Image in [Diffusion] Port the Wan VAE decoder fast paths to the Qwen-Image VAE #38020.Speed Tests and Profiling
Standalone VAE decode (torch.profiler GPU kernel self time, 81f 480x832):
End-to-end, TurboWan2.1-T2V-1.3B, 4 steps, 81f 480x832, idle H200,
mainand this PR run interleaved (mainfrom agit worktreeof the base commit), 6 prompts per run, median of the 5 warm requests:DecodingStage, losslessDecodingStage, highDmdDenoisingStage, losslessDmdDenoisingStage, highDenoising is unchanged, so the effect is confined to the VAE. With decode at ~60% of a warm request,
highsaves ~5% of the request and lossless ~1.4%.Qwen-Image also picks up the Triton upsample on its
highpath: 44.6 -> 42.9 ms at 1024x1024.Checklist
test_model_fast_paths.py+test_layout.py: 195 passed, 1 skipped;test_import_surface.py: 9 passed)Review and Merge Process
/tag-and-rerun-ci,/tag-run-ci-label,/rerun-failed-ciCI States
Latest PR Test (Base): 🚫 Run #34174270196
Latest PR Test (Extra): ❌ Run #34174270292
Latest PR Test (AMD ROCm 7.2): ❌ Run #34174270249