vllm_omni.diffusion.distributed.autoencoders.wan_decoder_fast_path ¶
Exact fast path for the Wan VAE decoder on a streaming (frame-at-a-time) decode.
A kernel ledger of the LingBot-World served path (4xH200, 480x832, width-sharded bf16 decode) put the decoder at 91 ms of GPU time per chunk per rank, half of it in ~2,300 launches of glue around the convolutions: autocast re-casting every conv weight and bias on every call (the autocast weight cache only serves leaves that require grad, so under torch.no_grad it never hits), the nearest upsample's fp32 round trip, and a freshly allocated, zero-filled, concatenated input tensor per conv call. None of that touches a value that reaches the convolution, so it can go without changing a single output bit:
conv_dtype: the decoder's convolution parameters are cast once to the dtype the decode already runs under (decode_autocast_dtype), the same cast autocast applied per call. Norm parameters stay in their own dtype, so the normalisation arithmetic and its fp32 promotion are untouched.WanUpsampleruns its nearest-exact gather on the activation dtype directly instead of viax.float()andtype_as: a gather moves values, so the result is identical.- The spatially sharded conv wrappers keep one input buffer per conv across calls (
reuse_input_buffer), writing only the activation interior and the halo slots into it.
Numerics: bit-identical to the plain path (tests/diffusion/distributed/test_wan_decoder_fast_path.py). Memory: one persistent input buffer per sharded conv (about the size of that conv's input).
Level "fused" adds the rest of what the ledger showed, at the cost of exactness (it is gated with the bf16 decode option, whose quality gate it shares): the decoder runs channels_last_3d end to end so cuDNN uses its NHWC kernels without the NCHW<->NHWC transposes it otherwise inserts around every convolution, every WanRMS_norm -> SiLU chain becomes one Triton kernel that reads and writes the activation dtype (the eager chain promotes to fp32 at * gamma and leaves the cast to the next convolution), the nearest upsample is a channels_last gather, and the upsample3d time-pair interleave writes channels_last_3d directly. Values differ from eager only through the norm's fp32 reduction order and the single rounding at the SiLU output.
FusedWanRMSNormSiLU ¶
Bases: Module
WanRMS_norm followed by SiLU as one kernel on channels_last_3d CUDA tensors.
Keeps the norm's parameters registered under their own names (...norm1.gamma), so state dicts and weight loading are unchanged. Off the kernel's domain (CPU, other layouts, grad) it runs the eager op chain, and a channels-first CUDA input is converted to channels_last_3d once (the attention block's gathered frame is the only such producer).
install_wan_decoder_fast_path ¶
install_wan_decoder_fast_path(
vae: Any,
*,
conv_dtype: dtype | None,
level: str = "exact",
) -> dict[str, int]
Install the fast path on vae's decoder (and post_quant_conv); idempotent.
conv_dtype must be the dtype the decode runs under (decode_autocast_dtype): with None the parameters are left alone, because casting them without autocast would change the convolution's arithmetic rather than remove a cast. level is "exact" or "fused" (see the module docstring).