Skip to content

WanCausalConv3d copies every activation with F.pad. Let the conv do the spatial padding. #14993

Description

@will-rice

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.

Activity

Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Metadata

Metadata

Assignees

No one assigned

    Labels

    No labels
    No labels

    Type

    No type

    Projects

    No projects

      Milestone

      No milestone

      Relationships

      None yet

      Development

      No branches or pull requests

      Issue actions