Skip to content

vllm_omni.diffusion.distributed.autoencoders.wan_decoder_kernels

Fused kernels for the Wan VAE decoder on channels_last_3d activations.

Adapted from SGLang's sglang.kernels.ops.diffusion (norm/wan_rmsnorm_silu_triton.py and layout/nearest_upsample_nhwc_triton.py, Apache-2.0), reduced to what the LingBot-World streaming decode uses. Each kernel comes with a can_use_* predicate and a torch reference of the eager op chain it replaces; the predicate decides, the kernel raises on an input it does not support.

  • :func:wan_rmsnorm_silu: SiLU(F.normalize(x, dim=1) * scale * gamma + bias) on a dense channels_last_3d [B, C, T, H, W] tensor, one program per pixel. fp32 channel statistics; the intermediate dtype boundaries of the eager chain are kept (normalise and * scale in x.dtype, * gamma in the promoted dtype, SiLU in fp32) and the result is stored in x.dtype, which is the dtype the following convolution reads under autocast anyway. Not bit-identical to aten (the channel reduction order differs), so it is gated with the bf16 decode option rather than mounted unconditionally.
  • :func:nearest_upsample_nhwc: integer-factor nearest upsample of a dense channels_last [N, C, H, W] tensor as a gather, bit-exact vs nn.Upsample in nearest and nearest-exact modes; aten's own NHWC nearest kernel is several times slower than its NCHW one, which is what a channels_last decoder would hit.

can_use_nearest_upsample_nhwc

can_use_nearest_upsample_nhwc(
    x: Tensor, scale_factor, mode: str
) -> bool

can_use_wan_rmsnorm_silu

can_use_wan_rmsnorm_silu(
    x: Tensor, gamma: Tensor, bias: Tensor | float | None
) -> bool

canonical_nhwc

canonical_nhwc(x: Tensor) -> bool

Dense channels_last with C > 1 and canonical strides on every dim, size-1 dims included.

nearest_upsample_nhwc

nearest_upsample_nhwc(x: Tensor, scale_factor) -> Tensor

Integer-factor nearest upsample of a canonical channels_last [N, C, H, W] tensor, dense channels_last out.

wan_rmsnorm_silu

wan_rmsnorm_silu(
    x: Tensor,
    gamma: Tensor,
    bias: Tensor | float | None,
    scale: float,
    eps: float = 1e-12,
    upcast: bool = True,
) -> Tensor

Fused SiLU(normalize(x) * scale * gamma + bias) on a dense channels_last_3d tensor; raises otherwise.

wan_rmsnorm_silu_reference

wan_rmsnorm_silu_reference(
    x: Tensor,
    gamma: Tensor,
    bias: Tensor | float,
    scale: float,
    eps: float = 1e-12,
    upcast: bool = True,
) -> Tensor

The eager norm chain followed by SiLU, cast back to x.dtype.

upcast=True is diffusers' WanRMS_norm (normalise in fp32, round to x.dtype); upcast=False is vLLM-Omni's RMSNormVAE (F.normalize on x itself, so every step rounds to x.dtype).