How to Use the WanDiffusionWrapper with Custom Model Configurations in LongLive
The WanDiffusionWrapper accepts a configuration dictionary (model_kwargs) that controls checkpoint loading, causal mode, sampling schedules, and VAE backends, allowing you to customize the Wan 5B diffusion pipeline without modifying core source code.
The WanDiffusionWrapper in the NVlabs/LongLive repository provides a high-level PyTorch interface for the Wan 5B video diffusion model. Rather than hardcoding hyperparameters, the wrapper uses a configuration-driven architecture where model paths, architecture sizes, and sampling schedules are defined in EasyDict objects. This design lets you swap checkpoints, enable streaming inference, or adjust diffusion schedules by simply modifying a Python dictionary before instantiation.
Understanding the WanDiffusionWrapper Architecture
The wrapper orchestrates multiple subsystems defined in utils/wan_5b_wrapper.py (lines 277–340). It handles the forward pass for conditional video generation while abstracting away the underlying transformer backbone, VAE, and noise scheduler.
Core Components
The wrapper integrates these key components:
- WanModel / CausalWanModel: The transformer diffusion backbone loaded via
from_pretrainedfromwan_5b/modules/model.pyorwan_5b/modules/causal_model.py. - FlowMatchScheduler: Provides diffusion timesteps and conversion utilities, defined in
utils/scheduler.py. - WanTextEncoder: Encodes text prompts into embeddings using the T5-based encoder (lines 16–55 of the wrapper file).
- WanVAEWrapper: Handles latent space encoding/decoding for video frames (lines 58–224).
Configuration-Driven Design
All hyperparameters reside in the wan_5b/configs/ directory. The base configuration in wan_5b/configs/shared_config.py defines defaults like checkpoint filenames, model dimensions, and patch sizes. Specific variants (e.g., wan_i2v_A14B.py) inherit from this base and override values for particular use cases.
Loading and Modifying Configurations
To customize the wrapper, clone an existing configuration and override specific fields before passing it as keyword arguments.
Cloning Built-in Configs
The recommended pattern uses copy.deepcopy to avoid mutating global defaults:
import copy
from wan_5b.configs.wan_i2v_A14B import i2v_A14B
custom_cfg = copy.deepcopy(i2v_A14B)
custom_cfg.model_name = "Wan2.2-TI2V-5B" # Directory under wan_models/
custom_cfg.timestep_shift = 6.0 # Controls diffusion speed
custom_cfg.is_causal = True # Enable streaming mode
custom_cfg.local_attn_size = 4 # Causal attention window
Key Configuration Parameters
When constructing WanDiffusionWrapper(**config), the constructor uses these critical fields:
model_name: Path to the checkpoint directory underwan_models/.is_causal: Boolean flag that selectsCausalWanModelinstead ofWanModelfor frame-by-frame generation.timestep_shift: Float value (default 8.0) passed toFlowMatchScheduler.set_timesteps().t_scale,rope_method,original_seq_len: Model-specific architectural attributes.
Instantiating the Wrapper with Custom Settings
After preparing your configuration dictionary, instantiate the wrapper and prepare input tensors. The following example demonstrates the complete setup:
from utils.wan_5b_wrapper import WanDiffusionWrapper, WanTextEncoder
import torch
# Initialize wrapper with custom configuration
generator = WanDiffusionWrapper(**custom_cfg)
# Prepare text conditioning
text_encoder = WanTextEncoder()
prompt = ["A sunrise over a mountain range."]
cond = {"prompt_embeds": text_encoder(prompt)["prompt_embeds"]}
# Create dummy noisy video input [B, C, T, H, W]
noisy = torch.randn(1, 3, 16, 256, 256).to(generator.device)
# Forward pass returns flow prediction and denoised output
flow_pred, pred_x0 = generator(
noisy,
cond,
timestep=torch.full((1, 1), 0.0)
)
The forward() method receives the noisy video tensor and a conditional_dict containing prompt_embeds. It optionally utilizes KV-cache for streaming inference, calls the underlying model via self._call_model, and converts flow-matching predictions to x0 predictions using _convert_flow_pred_to_x0.
Integrating Text Encoders and VAE Backends
The wrapper delegates video encoding to pluggable VAE wrappers selected via the build_vae_5b helper function (lines 505–527).
Switching VAE Types
You can swap the VAE backend by specifying vae_type in your arguments:
from utils.wan_5b_wrapper import build_vae_5b
from argparse import Namespace
args = Namespace(vae_type="mg_lightvae") # Alternative: "mg_lightvae_v2"
vae = build_vae_5b(args)
# Encode raw frames to latent space
frames = torch.randn(1, 3, 16, 256, 256)
latent = vae.encode_to_latent(frames)
# Decode back to pixel space
recon = vae.decode_to_pixel(latent)
Running Inference with Custom Configurations
For high-level training or inference, the DiffusionModel class in model/diffusion.py internally constructs a WanDiffusionWrapper from a model_kwargs dictionary:
from model.diffusion import DiffusionModel
# Assumes args contains model_kwargs with your custom config
diffusion = DiffusionModel(args)
# Generate video from text prompt
generated = diffusion.sample(
video_tensor,
prompt="A futuristic city at night"
)
Summary
- The
WanDiffusionWrapperinutils/wan_5b_wrapper.pyencapsulates the Wan 5B diffusion pipeline. - Configuration dictionaries control model selection (
is_causal), checkpoint paths (model_name), and sampling behavior (timestep_shift). - Clone existing configs from
wan_5b/configs/to customize without editing source code. - The wrapper requires
prompt_embedsfromWanTextEncoderand supports pluggable VAE backends viabuild_vae_5b. - Use
FlowMatchSchedulerparameters to adjust diffusion speed and noise characteristics.
Frequently Asked Questions
What file contains the WanDiffusionWrapper class definition?
The class is defined in utils/wan_5b_wrapper.py between lines 277 and 340 according to the LongLive source code. This file also contains the WanTextEncoder, VAE wrappers, and the build_vae_5b factory function.
How do I enable causal/streaming mode in WanDiffusionWrapper?
Set is_causal=True in your configuration dictionary before instantiation. This flag causes the wrapper to load CausalWanModel from wan_5b/modules/causal_model.py instead of the standard WanModel, enabling frame-by-frame streaming generation with KV-cache support.
Can I use a different VAE backend with WanDiffusionWrapper?
Yes. Pass a vae_type parameter (e.g., "mg_lightvae" or "mg_lightvae_v2") to the build_vae_5b helper function defined at lines 505–527 of utils/wan_5b_wrapper.py. This instantiates alternative VAE implementations like LightVAE5BWrapper without modifying the diffusion wrapper itself.
Where are the default configuration values defined?
Default values reside in wan_5b/configs/shared_config.py as an EasyDict named wan_shared_cfg. Specific model variants such as wan_i2v_A14B.py import and override these defaults to define concrete architectures and checkpoint paths.
Have a question about this repo?
These articles cover the highlights, but your codebase questions are specific. Give your agent direct access to the source. Share this with your agent to get started:
curl -s "https://instagit.com/install.md" Maintain an open-source project? Get it listed too →