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* scaleinx.dtype,* gammain the promoted dtype, SiLU in fp32) and the result is stored inx.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 vsnn.Upsampleinnearestandnearest-exactmodes; 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_wan_rmsnorm_silu ¶
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 ¶
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).