WanAnimate: Move face adapter blocks to corresponding main blocks to include them in block swap

This commit is contained in:
kijai
2025-09-21 16:34:17 +03:00
parent 0482667c78
commit fc3a684b0b
4 changed files with 36 additions and 28 deletions
+21
View File
@@ -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):
+1
View File
@@ -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()
+14 -13
View File
@@ -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
@@ -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__()