Skip to content

vllm_omni.diffusion.distributed.autoencoders.wan_vae_fastpath

Fast paths for the diffusers Wan (2.1/2.2) causal video VAE decoder.

install_wan_vae_fastpath(vae, level=...) rebinds the forwards of one loaded AutoencoderKLWan instance:

  • "lossless" (default): bit-exact rewrites of the decoder's data movement and normalization (fused Triton kernels with exact PyTorch fallbacks).
  • "channels_last": additionally converts decoder convolution weights to channels-last memory format (faster cuDNN kernels, not bit-exact).
  • "off": leave the diffusers implementation untouched.

The framework installs it from vllm_omni.diffusion.registry.initialize_model according to OmniDiffusionConfig.vae_fast_path.

Modules:

Name Description
decode

Chunked Wan decode with a preallocated output buffer.

forwards

Bit-exact replacement forwards for the diffusers Wan VAE decoder modules.

install

Instance-level installer for the Wan VAE decoder fast path.

triton_data_movement

Bit-exact data-movement kernels for the diffusers Wan causal VAE decoder.

triton_rms_norm

Bit-exact fused RMSNorm epilogue for the diffusers Wan VAE decoder.

triton_rms_norm_cl

Single-pass channels-last RMSNorm(+SiLU) for the Wan VAE decoder (Tier 2, not bit-exact).

triton_upsample

Bit-exact nearest-neighbour 2x spatial upsampling for the Wan VAE decoder.

REPORT_ATTR module-attribute

REPORT_ATTR = '_vllm_omni_wan_fastpath_report'

VAE_FAST_PATH_LEVELS module-attribute

VAE_FAST_PATH_LEVELS: tuple[str, ...] = (
    "off",
    "lossless",
    "channels_last",
)

WanVaeFastPathReport dataclass

What :func:install_wan_vae_fastpath did to one VAE instance.

channels_last class-attribute instance-attribute

channels_last: bool = False

fused_silu_dtypes class-attribute instance-attribute

fused_silu_dtypes: tuple[str, ...] = ()

installed instance-attribute

installed: bool

level instance-attribute

level: str

patched class-attribute instance-attribute

patched: Mapping[str, int] = field(default_factory=dict)

reason class-attribute instance-attribute

reason: str | None = None

decode_frames

decode_frames(vae, z: Tensor) -> Tensor

AutoencoderKLWan._decode's frame loop, writing every chunk straight into the result.

Upstream grows the result with torch.cat([out, out_], 2) on every chunk (quadratic copying) and then materializes unpatchify and clamp as two more full-size copies, so three or four copies of the decoded video are live at once. Here the final [B, C, T, H, W] buffer is allocated after the first chunk and each chunk is unpatchified and clamped directly into its slot. Every step is a permutation or an elementwise clamp, so the values are identical to upstream. Tiling dispatch is the caller's responsibility.

install_wan_vae_fastpath

install_wan_vae_fastpath(
    vae: Module, *, level: str = "lossless"
) -> WanVaeFastPathReport

Install the Wan decoder fast path on one loaded AutoencoderKLWan. Idempotent.

level: * "off": do nothing. * "lossless": bind the bit-exact replacement forwards (Tier 1). * "channels_last": Tier 1 plus channels-last conv weights (not bit-exact).

The installer refuses (and reports why) when the VAE is not a diffusers Wan VAE, or when spatially-sharded decode is configured or already installed: wan_spatial_shard rebinds the same forwards and replaces the causal convolutions with halo-exchanging variants.

is_installed

is_installed(vae: Module) -> bool

uninstall_wan_vae_fastpath

uninstall_wan_vae_fastpath(vae: Module) -> None

Restore original forwards and tensor layouts, retaining current weight values.