WanAnimate: Move face adapter blocks to corresponding main blocks to include them in block swap
This commit is contained in:
@@ -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):
|
||||
|
||||
@@ -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
@@ -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__()
|
||||
|
||||
Reference in New Issue
Block a user