wan_animate ¶
Wan2.2-Animate-14B: character animation / replacement.
Ported against both references: the official Wan-Video/Wan2.2 (wan/modules/animate/) and diffusers' merged WanAnimateTransformer3DModel. The official Diffusers-format checkpoint (Wan-AI/Wan2.2-Animate-14B-Diffusers) uses the diffusers naming, which this port loads directly.
Animate required no new tower: it is the stock Wan-I2V transformer (channel-concat conditioning, 36-channel input, CLIP image branch, single global timestep, standard RoPE grid), so this class subclasses WanTransformer3DModel -- the base __init__ builds the entire tower from WanAnimateArchConfig -- and adds three things:
pose_patch_embedding: a second patchifier whose output is added to the patchified video tokens, skipping the reference latent frame (frame 0 of the sequence is the reference image and carries no pose). Adding pose to frame 0 corrupts identity conditioning -- the off-by-one is load-bearing.- The face-driving stack (
wan_animate_face.py): 512x512 face crops -> LIA motion vectors -> causal 4x funnel -> 4+1 face tokens per latent frame. face_adapter: per-frame cross-attention applied residually after everyinject_face_latents_blocks-th block (blocks 0, 5, ..., 35 -- the adapter index i serves block i * stride; asserted in the arch config).
Sequence parallelism is not wired up in this first version: the face adapter's per-frame reshape needs the full token sequence, and an SP shard splits frames across ranks mid-frame. Single-GPU and FSDP-style weight sharding work (__init__ raises on sp_world_size > 1).
Classes¶
fastvideo.models.dits.wan_animate.WanAnimateTransformer3DModel ¶
Bases: WanTransformer3DModel
Wan2.2-Animate-14B transformer.
Source code in fastvideo/models/dits/wan_animate.py
Methods:¶
fastvideo.models.dits.wan_animate.WanAnimateTransformer3DModel.forward ¶
forward(hidden_states: Tensor, encoder_hidden_states: Tensor | list[Tensor], timestep: LongTensor, encoder_hidden_states_image: Tensor | list[Tensor] | None = None, pose_latents: Tensor | None = None, face_pixel_values: Tensor | None = None, guidance=None, **kwargs) -> Tensor
Denoise one step of an animation/replacement segment.
hidden_states [B, 36, T+1, H, W]: noise | 4ch mask | conditional latent y, channel-concatenated by the pipeline. Frame 0 is the reference latent slot. pose_latents [B, 16, T, H, W] VAE-encoded skeleton video; exactly one frame fewer than hidden_states (no pose on ref). face_pixel_values [B, 3, F, 512, 512] raw face crops, F = 4T - 3 pixel frames for T latent frames. encoder_hidden_states_image CLIP features of the reference image (257 tokens) -- required, the I2V cross-attention splits them off positionally.
Returns the full [B, 16, T+1, H, W] velocity; the decode stage drops the reference slot (and any guidance frames).
Source code in fastvideo/models/dits/wan_animate.py
159 160 161 162 163 164 165 166 167 168 169 170 171 172 173 174 175 176 177 178 179 180 181 182 183 184 185 186 187 188 189 190 191 192 193 194 195 196 197 198 199 200 201 202 203 204 205 206 207 208 209 210 211 212 213 214 215 216 217 218 219 220 221 222 223 224 225 226 227 228 229 230 231 232 233 234 235 236 237 238 239 240 241 242 243 244 245 246 247 248 249 250 251 252 253 254 255 256 | |
fastvideo.models.dits.wan_animate.WanAnimateTransformer3DModel.materialize_non_persistent_buffers ¶
Rebuild the motion encoder's anti-aliasing blur kernels after meta load.
TransformerLoader builds the model under torch.device("meta") and streams checkpoint weights in; non-persistent buffers are not in the checkpoint and stay on meta until this hook (called by fsdp_load) recreates them. Kernels stay fp32 -- forward casts them to the activation dtype.