From fc3a684b0b42fe701c2d1c7b5b5744456519090d Mon Sep 17 00:00:00 2001 From: kijai <40791699+kijai@users.noreply.github.com> Date: Sun, 21 Sep 2025 16:34:17 +0300 Subject: [PATCH] WanAnimate: Move face adapter blocks to corresponding main blocks to include them in block swap --- nodes_model_loading.py | 21 +++++++++++++++++ nodes_sampler.py | 1 + wanvideo/modules/model.py | 27 +++++++++++----------- wanvideo/modules/wananimate/face_blocks.py | 15 ------------ 4 files changed, 36 insertions(+), 28 deletions(-) diff --git a/nodes_model_loading.py b/nodes_model_loading.py index ac57738..c82b88a 100644 --- a/nodes_model_loading.py +++ b/nodes_model_loading.py @@ -4,6 +4,7 @@ import os, gc, uuid from .utils import log, apply_lora import numpy as np from tqdm import tqdm +import re from .wanvideo.modules.model import WanModel, LoRALinearLayer from .wanvideo.modules.t5 import T5EncoderModel @@ -758,6 +759,17 @@ class WanVideoSetLoRAs: return (patcher,) +def rename_fuser_block(name): + # map fuser blocks to main blocks + new_name = name + if "face_adapter.fuser_blocks." in name: + match = re.search(r'face_adapter\.fuser_blocks\.(\d+)\.', name) + if match: + fuser_block_num = int(match.group(1)) + main_block_num = fuser_block_num * 5 + new_name = name.replace(f"face_adapter.fuser_blocks.{fuser_block_num}.", f"blocks.{main_block_num}.fuser_block.") + return new_name + def load_weights(transformer, sd=None, weight_dtype=None, base_dtype=None, transformer_load_device=None, block_swap_args=None, gguf=False, reader=None, patcher=None): params_to_keep = {"time_in", "patch_embedding", "time_", "modulation", "text_embedding", @@ -784,6 +796,7 @@ def load_weights(transformer, sd=None, weight_dtype=None, base_dtype=None, all_tensors.extend(r.tensors) for tensor in all_tensors: name = tensor.name + name = rename_fuser_block(name) if "glob" not in name and "audio_proj" in name: name = name.replace("audio_proj", "multitalk_audio_proj") load_device = device @@ -1052,6 +1065,14 @@ class WanVideoModelLoader: sd, reader = load_gguf(model_path) gguf_reader.append(reader) + is_wananimate = "pose_patch_embedding.weight" in sd + # rename WanAnimate face fuser block keys to insert into main blocks instead + if is_wananimate: + for key in list(sd.keys()): + new_key = rename_fuser_block(key) + if new_key != key: + sd[new_key] = sd.pop(key) + if quantization == "disabled": for k, v in sd.items(): if isinstance(v, torch.Tensor): diff --git a/nodes_sampler.py b/nodes_sampler.py index 26c4c36..8244626 100644 --- a/nodes_sampler.py +++ b/nodes_sampler.py @@ -2644,6 +2644,7 @@ class WanVideoSampler: cm_result = cm.transfer(src=img, ref=ref_images.permute(1, 2, 3, 0).squeeze(0).cpu().float().numpy(), method=colormatch) cm_result_list.append(torch.from_numpy(cm_result).to(vae.dtype)) videos = torch.stack(cm_result_list, dim=0).permute(3, 0, 1, 2) + del cm_result_list current_ref_images = videos[:, -1:].clone().detach() diff --git a/wanvideo/modules/model.py b/wanvideo/modules/model.py index deab0e5..db46bfb 100644 --- a/wanvideo/modules/model.py +++ b/wanvideo/modules/model.py @@ -838,6 +838,7 @@ class WanAttentionBlock(nn.Module): rope_func="comfy", use_motion_attn=False, use_humo_audio_attn=False, + face_fuser_block=False ): super().__init__() self.dim = out_features @@ -855,6 +856,7 @@ class WanAttentionBlock(nn.Module): self.kv_cache = None self.use_motion_attn = use_motion_attn + self.has_face_fuser_block = face_fuser_block # layers self.norm1 = WanLayerNorm(out_features, eps) @@ -886,6 +888,10 @@ class WanAttentionBlock(nn.Module): if use_humo_audio_attn: self.audio_cross_attn_wrapper = AudioCrossAttentionWrapper(in_features, out_features, num_heads, qk_norm, eps, kv_dim=1536) + if face_fuser_block: + from .wananimate.face_blocks import FaceBlock + self.fuser_block = FaceBlock(self.dim, num_heads) + #@torch.compiler.disable() def get_mod(self, e): if e.dim() == 3: @@ -1661,7 +1667,8 @@ class WanModel(torch.nn.Module): self.blocks = nn.ModuleList([ WanAttentionBlock(cross_attn_type, self.in_features, self.out_features, ffn_dim, ffn2_dim, num_heads, qk_norm, cross_attn_norm, eps, - attention_mode=self.attention_mode, rope_func=self.rope_func, use_motion_attn=(i % 4 == 0 and use_motion_attn), use_humo_audio_attn=self.humo_audio) + attention_mode=self.attention_mode, rope_func=self.rope_func, use_motion_attn=(i % 4 == 0 and use_motion_attn), use_humo_audio_attn=self.humo_audio, + face_fuser_block = i % 5 == 0 and is_wananimate) for i in range(num_layers) ]) #MTV Crafter @@ -1752,14 +1759,10 @@ class WanModel(torch.nn.Module): self.motion_encoder = self.pose_patch_embedding = self.face_encoder = self.face_adapter = None if is_wananimate: from .wananimate.motion_encoder import MotionExtractor - from .wananimate.face_blocks import FaceEncoder, FaceAdapter + from .wananimate.face_blocks import FaceEncoder self.pose_patch_embedding = nn.Conv3d(16, dim, kernel_size=patch_size, stride=patch_size) self.motion_encoder = MotionExtractor() - self.face_adapter = FaceAdapter( - num_heads=self.num_heads, - feature_dim=self.dim, - num_adapter_layers=self.num_layers // 5, - ) + self.face_encoder = FaceEncoder( in_dim=motion_encoder_dim, out_dim=self.dim, @@ -1931,12 +1934,10 @@ class WanModel(torch.nn.Module): return torch.cat([pad_face, motion_vec], dim=1) - def wananimate_forward(self, block_idx, x, motion_vec, strength=1.0, motion_masks=None): - if block_idx % 5 == 0: + def wananimate_forward(self, block, x, motion_vec, strength=1.0, motion_masks=None): adapter_args = [x, motion_vec, motion_masks] - residual_out = self.face_adapter.fuser_blocks[block_idx // 5](*adapter_args) + residual_out = block.fuser_block(*adapter_args) return x.add(residual_out, alpha=strength) - return x def rope_encode_comfy(self, t, h, w, freq_offset=0, t_start=0, attn_cond_shape=None, steps_t=None, steps_h=None, steps_w=None, ntk_alphas=[1,1,1], device=None, dtype=None): @@ -2649,8 +2650,8 @@ class WanModel(torch.nn.Module): x, x_ip = block(x, x_ip=x_ip, **kwargs) #run block if self.audio_injector is not None and s2v_audio_input is not None: x = self.audio_injector_forward(b, x, merged_audio_emb, scale=s2v_audio_scale) #s2v - if self.motion_encoder is not None and motion_vec is not None: - x = self.wananimate_forward(b, x, motion_vec, strength=wananim_face_strength) + if block.has_face_fuser_block and motion_vec is not None: + x = self.wananimate_forward(block, x, motion_vec, strength=wananim_face_strength) if self.block_swap_debug: compute_end = time.perf_counter() compute_time = compute_end - compute_start diff --git a/wanvideo/modules/wananimate/face_blocks.py b/wanvideo/modules/wananimate/face_blocks.py index ead9856..e5399af 100644 --- a/wanvideo/modules/wananimate/face_blocks.py +++ b/wanvideo/modules/wananimate/face_blocks.py @@ -82,21 +82,6 @@ class RMSNorm(nn.Module): return output -class FaceAdapter(nn.Module): - def __init__(self, feature_dim, num_heads, num_adapter_layers=1, dtype=None, device=None): - super().__init__() - self.fuser_blocks = nn.ModuleList([FaceBlock(feature_dim, num_heads, device=device, dtype=dtype) for _ in range(num_adapter_layers)]) - - def forward( - self, - x: torch.Tensor, - motion_embed: torch.Tensor, - idx: int, - ) -> torch.Tensor: - - return self.fuser_blocks[idx](x, motion_embed) - - class FaceBlock(nn.Module): def __init__(self, feature_dim, num_heads, dtype=None, device=None): super().__init__()