Is your feature request related to a problem? Please describe.
WanCausalConv3d.forward zero-pads height, width and time with F.pad, then runs an unpadded nn.Conv3d. F.pad writes a full copy of the activation on every call, and the Wan VAE calls this for every conv on every latent frame.
Describe the solution you'd like.
Only the time axis needs manual padding, because it is causal and the amount depends on the cached frames. Height and width are plain symmetric zero padding, so the conv can do them itself:
self.temporal_padding = 2 * self.padding[0]
self.padding = (0, self.padding[1], self.padding[2])
forward then pads only the time axis, and skips F.pad when the cache already supplies the temporal context. The result is the same, since the module uses zero padding.
On an A10G with torch 2.11.0+cu128, AutoencoderKLWan.decode for 33 frames at 512x512 gets about 6% faster:
| dtype |
mode |
main |
with the change |
| float32 |
eager |
5150 ms |
4838 ms |
| float32 |
torch.compile |
4888 ms |
4579 ms |
| float16 |
eager |
3048 ms |
2873 ms |
| float16 |
torch.compile |
2727 ms |
2553 ms |
| bfloat16 |
eager |
3046 ms |
2872 ms |
| bfloat16 |
torch.compile |
2726 ms |
2552 ms |
Peak memory is unchanged. Decoding the latents of a real 33-frame 512x512 video, the output is bit-identical to main on CPU (float32 and float64) and on GPU with cuDNN disabled. With cuDNN, 0.3% of 8-bit values change by one level (max difference 1.6e-3 on a [-1, 1] scale, 94.9 dB PSNR), because cuDNN picks its algorithm from the call's arguments.
I would cover two classes in one PR:
WanCausalConv3d in autoencoder_kl_wan.py
QwenImageCausalConv3d in autoencoder_kl_qwenimage.py, which is a hand copy of it
Describe alternatives you've considered.
QwenImage21CausalConv3d has the 2D version of the same pattern (F.pad over height and width, then an unpadded nn.Conv2d). I have not measured it, so I would leave it out unless you want it in the same PR.
Additional context.
The change replaces the private _padding attribute with temporal_padding. Nothing else in src/ reads _padding, and state-dict keys do not change.
I have the fix ready on a branch, with a test that checks the new conv against the original pad-then-conv math for each kernel configuration, with and without cached frames. I used an AI agent to help prepare it, so per the contributing guide I am opening this issue first. I can open the PR as soon as a maintainer is happy with the approach.
Is your feature request related to a problem? Please describe.
WanCausalConv3d.forwardzero-pads height, width and time withF.pad, then runs an unpaddednn.Conv3d.F.padwrites a full copy of the activation on every call, and the Wan VAE calls this for every conv on every latent frame.Describe the solution you'd like.
Only the time axis needs manual padding, because it is causal and the amount depends on the cached frames. Height and width are plain symmetric zero padding, so the conv can do them itself:
forwardthen pads only the time axis, and skipsF.padwhen the cache already supplies the temporal context. The result is the same, since the module uses zero padding.On an A10G with torch 2.11.0+cu128,
AutoencoderKLWan.decodefor 33 frames at 512x512 gets about 6% faster:maintorch.compiletorch.compiletorch.compilePeak memory is unchanged. Decoding the latents of a real 33-frame 512x512 video, the output is bit-identical to
mainon CPU (float32 and float64) and on GPU with cuDNN disabled. With cuDNN, 0.3% of 8-bit values change by one level (max difference 1.6e-3 on a [-1, 1] scale, 94.9 dB PSNR), because cuDNN picks its algorithm from the call's arguments.I would cover two classes in one PR:
WanCausalConv3dinautoencoder_kl_wan.pyQwenImageCausalConv3dinautoencoder_kl_qwenimage.py, which is a hand copy of itDescribe alternatives you've considered.
QwenImage21CausalConv3dhas the 2D version of the same pattern (F.padover height and width, then an unpaddednn.Conv2d). I have not measured it, so I would leave it out unless you want it in the same PR.Additional context.
The change replaces the private
_paddingattribute withtemporal_padding. Nothing else insrc/reads_padding, and state-dict keys do not change.I have the fix ready on a branch, with a test that checks the new conv against the original pad-then-conv math for each kernel configuration, with and without cached frames. I used an AI agent to help prepare it, so per the contributing guide I am opening this issue first. I can open the PR as soon as a maintainer is happy with the approach.