Repository navigation
[Diffusion] Fuse the Wan VAE conv bias epilogue and tile the RMSNorm+SiLU kernel - #38650
Open
Dayuxiaoshui wants to merge 1 commit into
Open
Dayuxiaoshui wants to merge 1 commit into
Dayuxiaoshui wants to merge 1 commit into
Conversation
…SiLU kernel Two costs in the Wan 2.1 VAE decoder that kernel-level profiles hide: 1. PyTorch's cuDNN conv path adds the bias as a separate broadcast add_ after cudnn_convolution, and aten does not vectorise that add on channels_last outputs (1.5 TB/s on H200; 797 launches, ~10% of the whole decode in both quality modes). It only shows up at the aten-op level as conv3d minus cudnn_convolution. 2. wan_rmsnorm_silu ran one pixel per program with 4 warps over 96-384 channels: 0.63 TB/s on [1, 96, 4, 480, 832] bf16, 16% of the high-quality decode. Lossless (bit-exact, unconditional): - New kernel modulate/conv_bias_epilogue_triton.py: x.dtype(x + bias) with aten's arithmetic (fp32 add, one rounding), optionally fused with the residual add as x.dtype(x.dtype(x + bias) + h), the two roundings of the eager conv(x) + h chain. 4.1 TB/s; output keeps the input's dense channels_last(_3d) strides like the in-place add_. - WanCausalConv3d._conv_with_epilogue runs the conv without bias and applies it through that kernel; residual_block_forward fuses conv2's bias with the block's residual add (residual=) and defers conv1's bias into norm2 (skip_bias= + WanRMS_norm.forward(x, conv_bias=...)), where the eager norm adds it first with the same arithmetic. Falls back to the exact aten ops. Quality-gated: - wan_rmsnorm_silu tiles [ROWS, C] pixels per program (32 rows at C <= 128, down to 4 at C = 1024): 0.63 -> ~3 TB/s. The 2D reduction tree moves ~2e-6 of the elements by one bf16 ulp relative to the previous kernel; both sit at the same distance from eager (20% of elements differ), so the gated contract is unchanged. Optional conv_bias= folds the deferred conv1 bias into the same pass, applied exactly like aten's add_ before the statistics, so the norm sees the values it would have read. - FusedWanRMSNormSiLU.forward(x, conv_bias=None) routes the deferred bias: folded under the gate, else applied by the bit-exact epilogue first. Wan 2.1 VAE, 81 frames 480x832, H200 (torch.profiler GPU kernel time): bf16 weights: lossless 1324 -> 1253 ms, high 1112 -> 935 ms; fp32 weights + bf16 autocast (the T2V pipeline setting): lossless 1593 -> 1530 ms, high 1215 -> 1056 ms. Qwen-Image high 43.3 -> 39.4 ms (shared norm kernel). Lossless torch.equal vs main in both settings; TurboWan2.1-T2V-1.3B end-to-end: 6/6 lossless mp4s md5-identical to main, warm DecodingStage 932 -> 785 ms at quality=high on a quiet GPU. Tests: test_conv_bias_epilogue.py (aten oracle, dtypes, 4D/5D, residual, predicate rejections, conv3d-with-bias equality); test_norm.py multi-row ragged tiles and exact conv_bias fold; test_model_fast_paths.py three-chunk tiny-decoder decode fused vs pure aten (torch.equal), deferred bias through the gated wrapper off (exact) and on (bounded).
Dayuxiaoshui
requested review from
AgainstEntropy,
BBuf,
HaiShaw,
mickqian,
ping1jing2,
yichiche and
yingluosanqian
as code owners
September 9, 2026 08:13
This was referenced Sep 11, 2026
This was referenced Sep 19, 2026
This was referenced Oct 8, 2026
This branch has not been deployed
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 and #38182. Two costs in the Wan 2.1 VAE decoder that kernel-level profiles hide, found by grouping a
quality=highdecode at the aten-op level (key_averages(group_by_input_shape=True)):cudnn_convolutionand thenoutput.add_(bias.view(1, C, 1, 1, 1)). On channels_last outputs aten does not vectorise that broadcast add: 1.5 TB/s on H200 for[1, 96, 4, 480, 832]bf16. It shows up only asaten::conv3dminusaten::cudnn_convolution: 797 launches, 278 ms of a 2847 ms lossless decode and 227 ms of a 1876 ms high decode on a shared GPU (~10% in both modes).wan_rmsnorm_siluran one pixel per program with 4 warps over 96-384 channels: 0.63 TB/s on the largest shape, 16% of the high-quality decode.Modifications
New kernel, lossless (bit-exact, unconditional)
kernels/ops/diffusion/modulate/conv_bias_epilogue_triton.py:x.dtype(float(x) + float(bias)), the exact arithmetic of aten'sadd_(fp32 opmath, one rounding), optionally fused with the residual add asx.dtype(x.dtype(x + bias) + h), i.e. the two roundings of the eagerconv(x) + hchain. 4.1 TB/s.biasmay be fp32 on a half-precisionx(autocast): it is cast tox.dtypefirst, which is what autocast does before the conv call. The output isempty_like(x), so it keeps the dense channels_last(_3d) strides like the in-placeadd_. Predicate requiresC > 1with canonical strides and never raises; the kernel validates and raisesValueError.WanCausalConv3d._conv_with_epilogue: runs the conv without bias and applies it through the kernel; falls back to the exact aten ops (out.add_(bias.to(out.dtype).view(...))).residual_block_forward: conv2's bias and the block's residual add become one pass (residual=h); conv1's bias is deferred into norm2 (skip_bias=True+WanRMS_norm.forward(x, conv_bias=...), where the eager norm adds it first with the same arithmetic).Quality-gated
wan_rmsnorm_silutiles[ROWS, C]consecutive pixels per program (32 rows atC <= 128down to 4 atC = 1024): 0.63 -> ~3 TB/s atC = 96. The 2D reduction tree moves ~2e-6 of the elements by one bf16 ulp relative to the previous kernel, while both differ from eager on ~20% of elements, so the gated contract is unchanged. Optionalconv_bias=folds the deferred conv1 bias into the same pass, applied exactly like aten'sadd_before the statistics.FusedWanRMSNormSiLU.forward(x, conv_bias=None): folded under the gate, otherwise applied by the bit-exact epilogue kernel before the eager chain.Tests / docs
test_conv_bias_epilogue.py: aten oracle (x.clone().add_(bias),+ h) for bf16/fp16/fp32 x fp32/half bias, 4D and 5D, stride identity;F.conv3dwith bias vs conv-without-bias + epilogue; predicate rejections (NCDHW, wrong bias width/dtype, residual layout/dtype, empty,C == 1,requires_grad).test_norm.py: ragged multi-row tiles at the three tile configurations;conv_biasfold equals the kernel on the pre-biased tensor bit for bit.test_model_fast_paths.py: a three-chunk tinyWanDecoder3ddecode through the feature cache, fused path vs pure aten (torch.equal); deferred bias through the gated wrapper with the gate off (exact) and on (bounded relative error).existing-fast-paths.mdupdated.Accuracy Tests
1x H200, Wan 2.1 VAE weights (
Wan-AI/Wan2.1-I2V-14B-480P-Diffusers), 81 frames 480x832.Lossless
torch.equalvsmain: True.torch.equalvsmain: True.sglang generate, TurboWan2.1-T2V-1.3B, 4 steps, seed 1, 6 prompts: all 6 mp4s md5-identical tomain(and to the outputs of the two previous PRs).Quality path
Standalone high vs
mainlossless on the same latent: PSNR 58.4 dB (bf16), 59.1 dB (fp32 + autocast), unchanged from before this PR; the multi-row norm kernel differs from the previous gated kernel on ~2e-6 of elements by one bf16 ulp.Speed Tests and Profiling
Standalone VAE decode, torch.profiler GPU kernel self time, idle H200:
Per-kernel (bf16, high):
_wan_rmsnorm_silu_kernel177 -> 57 ms over 609 launches; the 797 aten biasadd_launches are replaced by 396_conv_bias_epilogue_kernellaunches totalling 22 ms (649 launches / 35 ms on the lossless path). Qwen-Image 1024x1024 high: 43.3 -> 39.4 ms (shared norm kernel), lossless unchanged and still bit-exact.End-to-end, TurboWan2.1-T2V-1.3B, 4 steps, 81f 480x832,
mainfrom agit worktreeof the base commit, 6 prompts, median of the 5 warm requests. The decode stage is ~60% of a warm request. This box is shared; runs whoseDmdDenoisingStage(identical code on both sides) drifted were discarded, and the cleanest interleaved pair is reported:DecodingStage, highDecodingStage, losslessDmdDenoisingStageMicro-benchmarks behind the design (idle H200, bf16
[1, 96, 4, 480, 832]): aten biasadd_0.409 ms (1.5 TB/s), Triton bias add 0.150 ms (4.1 TB/s), atenadd_+ residualadd0.621 ms vs fused 0.215 ms; norm 0.970 ms -> 0.201 ms.Checklist
test_conv_bias_epilogue.py,test_norm.py,test_layout.py,test_model_fast_paths.py: 393 passed, 1 skipped)Review and Merge Process
/tag-and-rerun-ci,/tag-run-ci-label,/rerun-failed-ciCI States
Latest PR Test (Base): ❌ Run #34327993512
Latest PR Test (Extra): ❌ Run #34327993130
Latest PR Test (AMD ROCm 7.2): ❌ Run #34327993453