init: support lite model

This commit is contained in:
kijai
2025-09-27 00:49:02 +03:00
parent 37365817e8
commit ab31158673
9 changed files with 712 additions and 25 deletions
+9
View File
@@ -61,6 +61,13 @@ except Exception as e:
HUMO_NODE_CLASS_MAPPINGS = {}
HUMO_NODE_DISPLAY_NAME_MAPPINGS = {}
try:
from .lynx.nodes import NODE_CLASS_MAPPINGS as LYNX_NODE_CLASS_MAPPINGS, NODE_DISPLAY_NAME_MAPPINGS as LYNX_NODE_DISPLAY_NAME_MAPPINGS
except Exception as e:
log.warning(f"WanVideoWrapper WARNING: Lynx nodes not available due to error in importing them: {e}")
LYNX_NODE_CLASS_MAPPINGS = {}
LYNX_NODE_DISPLAY_NAME_MAPPINGS = {}
NODE_CLASS_MAPPINGS.update(RECAM_MASTER_NODE_CLASS_MAPPINGS)
NODE_CLASS_MAPPINGS.update(UNIANIMATE_NODE_CLASS_MAPPINGS)
NODE_CLASS_MAPPINGS.update(SKYREELS_NODE_CLASS_MAPPINGS)
@@ -80,6 +87,7 @@ NODE_CLASS_MAPPINGS.update(MTV_NODE_CLASS_MAPPINGS)
NODE_CLASS_MAPPINGS.update(S2V_NODE_CLASS_MAPPINGS)
NODE_CLASS_MAPPINGS.update(HUMO_NODE_CLASS_MAPPINGS)
NODE_CLASS_MAPPINGS.update(SAMPLER_NODE_CLASS_MAPPINGS)
NODE_CLASS_MAPPINGS.update(LYNX_NODE_CLASS_MAPPINGS)
NODE_DISPLAY_NAME_MAPPINGS.update(RECAM_MASTER_NODE_DISPLAY_NAME_MAPPINGS)
NODE_DISPLAY_NAME_MAPPINGS.update(UNIANIMATE_NODE_DISPLAY_NAME_MAPPINGS)
@@ -100,5 +108,6 @@ NODE_DISPLAY_NAME_MAPPINGS.update(MTV_NODE_DISPLAY_NAME_MAPPINGS)
NODE_DISPLAY_NAME_MAPPINGS.update(S2V_NODE_DISPLAY_NAME_MAPPINGS)
NODE_DISPLAY_NAME_MAPPINGS.update(HUMO_NODE_DISPLAY_NAME_MAPPINGS)
NODE_DISPLAY_NAME_MAPPINGS.update(SAMPLER_NODE_DISPLAY_NAME_MAPPINGS)
NODE_DISPLAY_NAME_MAPPINGS.update(LYNX_NODE_DISPLAY_NAME_MAPPINGS)
__all__ = ["NODE_CLASS_MAPPINGS", "NODE_DISPLAY_NAME_MAPPINGS"]
+114
View File
@@ -0,0 +1,114 @@
# Copyright 2025 Bytedance Ltd. and/or its affiliates
# SPDX-License-Identifier: Apache-2.0
import torch
import torchvision.transforms as T
import numpy as np
from insightface.utils import face_align
from insightface.app import FaceAnalysis
#from facexlib.recognition import init_recognition_model
__all__ = [
"FaceEncoderArcFace",
"get_landmarks_from_image",
]
detector = None
def get_landmarks_from_image(image):
"""
Detect landmarks with insightface.
Args:
image (np.ndarray or PIL.Image):
The input image in RGB format.
Returns:
5 2D keypoints, only one face will be returned.
"""
global detector
if detector is None:
detector = FaceAnalysis()
detector.prepare(ctx_id=0, det_size=(640, 640))
in_image = np.array(image).copy()
faces = detector.get(in_image)
if len(faces) == 0:
raise ValueError("No face detected in the image")
# Get the largest face
face = max(faces, key=lambda x: (x.bbox[2] - x.bbox[0]) * (x.bbox[3] - x.bbox[1]))
# Return the 5 keypoints directly
keypoints = face.kps # 5 x 2
return keypoints
from facexlib.utils import load_file_from_url
from facexlib.recognition.arcface_arch import Backbone
def init_recognition_model(model_name, half=False, device='cuda', model_rootpath=None):
print("Initializing recognition model:", model_name)
if model_name == 'arcface':
model = Backbone(num_layers=50, drop_ratio=0.6, mode='ir_se').to('cuda').eval()
model_url = 'https://github.com/xinntao/facexlib/releases/download/v0.1.0/recognition_arcface_ir_se50.pth'
else:
raise NotImplementedError(f'{model_name} is not implemented.')
model_path = load_file_from_url(
url=model_url, model_dir='facexlib/weights', progress=True, file_name=None, save_dir=model_rootpath)
print("Loading model from:", model_path)
model.load_state_dict(torch.load(model_path), strict=True)
model.eval()
model = model.to(device)
return model
class FaceEncoderArcFace():
""" Official ArcFace, no_grad-only """
def __repr__(self):
return "ArcFace"
def init_encoder_model(self, device, eval_mode=True):
self.device = device
self.encoder_model = init_recognition_model('arcface', device=device)
if eval_mode:
self.encoder_model.eval()
@torch.no_grad()
def input_preprocessing(self, in_image, landmarks, image_size=112):
assert landmarks is not None, "landmarks are not provided!"
in_image = np.array(in_image)
landmark = np.array(landmarks)
face_aligned = face_align.norm_crop(in_image, landmark=landmark, image_size=image_size)
image_transform = T.Compose([
T.ToTensor(),
T.Normalize([0.5], [0.5]),
])
face_aligned = image_transform(face_aligned).unsqueeze(0).to(self.device)
return face_aligned
@torch.no_grad()
def __call__(self, in_image, need_proc=False, landmarks=None, image_size=112):
if need_proc:
in_image = self.input_preprocessing(in_image, landmarks, image_size)
else:
assert isinstance(in_image, torch.Tensor)
image_embeds = self.encoder_model(in_image[:, [2, 1, 0], :, :].contiguous()) # [B, 512], normalized
return image_embeds, in_image
+63
View File
@@ -0,0 +1,63 @@
# Copyright (c) 2022 Insightface Team
# Copyright (c) 2025 Bytedance Ltd. and/or its affiliates
# This file has been modified by Bytedance Ltd. and/or its affiliates on September 15, 2025.
# SPDX-License-Identifier: Apache-2.0
# Original file (insightface) was released under MIT License:
# Permission is hereby granted, free of charge, to any person obtaining a copy
# of this software and associated documentation files (the "Software"), to deal
# in the Software without restriction, including without limitation the rights
# to use, copy, modify, merge, publish, distribute, sublicense, and/or sell
# copies of the Software, and to permit persons to whom the Software is
# furnished to do so, subject to the following conditions:
# The above copyright notice and this permission notice shall be included in all
# copies or substantial portions of the Software.
# THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR
# IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY,
# FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE
# AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER
# LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM,
# OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN THE
# SOFTWARE.
import numpy as np
import cv2
from PIL import Image
def estimate_norm(lmk, image_size=112, arcface_dst=None):
from skimage import transform as trans
assert lmk.shape == (5, 2)
assert image_size%112==0 or image_size%128==0
if image_size%112==0:
ratio = float(image_size)/112.0
diff_x = 0
else:
ratio = float(image_size)/128.0
diff_x = 8.0*ratio
dst = arcface_dst * ratio
dst[:,0] += diff_x
tform = trans.SimilarityTransform()
tform.estimate(lmk, dst)
M = tform.params[0:2, :]
return M
def get_arcface_dst(extend_face_crop=False, extend_ratio=0.8):
arcface_dst = np.array(
[[38.2946, 51.6963], [73.5318, 51.5014], [56.0252, 71.7366],
[41.5493, 92.3655], [70.7299, 92.2041]], dtype=np.float32)
if extend_face_crop:
arcface_dst[:,1] = arcface_dst[:,1] + 10
arcface_dst = (arcface_dst - 112/2) * extend_ratio + 112/2
return arcface_dst
def align_face(image_pil, face_kpts, extend_face_crop=False, extend_ratio=0.8, face_size=112):
arcface_dst = get_arcface_dst(extend_face_crop, extend_ratio)
M = estimate_norm(face_kpts, face_size, arcface_dst)
image_cv2 = cv2.cvtColor(np.array(image_pil), cv2.COLOR_RGB2BGR)
face_image_cv2 = cv2.warpAffine(image_cv2, M, (face_size, face_size), borderValue=0.0)
face_image = Image.fromarray(cv2.cvtColor(face_image_cv2, cv2.COLOR_BGR2RGB))
return face_image
+75
View File
@@ -0,0 +1,75 @@
import torch
import torch.nn as nn
from typing import Optional, List
from ..wanvideo.modules.attention import attention
def vector_to_list(tensor, lens, dim):
return list(torch.split(tensor, lens, dim=dim))
def list_to_vector(tensor_list, dim):
lens = [tensor.shape[dim] for tensor in tensor_list]
tensor = torch.cat(tensor_list, dim)
return tensor, lens
def merge_token_lists(list1, list2, dim):
assert(len(list1) == len(list2))
return [torch.cat((t1, t2), dim) for t1, t2 in zip(list1, list2)]
class WanLynxIPCrossAttention(nn.Module):
def __init__(self, cross_attention_dim=5120, dim=5120, n_registers=16, bias=True):
super().__init__()
self.to_k_ip = nn.Linear(cross_attention_dim, dim, bias=bias)
self.to_v_ip = nn.Linear(cross_attention_dim, dim, bias=bias)
if n_registers > 0:
self.registers = nn.Parameter(torch.randn(1, n_registers, cross_attention_dim) / dim**0.5)
else:
self.registers = None
def forward(self, block, q, x, ip_x):
b, n, d = x.size(0), block.num_heads, block.head_dim
print("ip_x.shape", ip_x.shape) #torch.Size([1, 16, 5120])
if self.registers is not None:
print("self.registers.shape", self.registers.shape) #torch.Size([1, 16, 5120])
print("ip_x.shape", ip_x.shape) #torch.Size([1, 16, 5120])
ip_lens = [ip_x.shape[1]]
ip_x_list = vector_to_list(ip_x, ip_lens, 1)
ip_x_list = merge_token_lists(ip_x_list, [self.registers] * len(ip_x_list), 1)
ip_x, ip_lens = list_to_vector(ip_x_list, 1)
ip_key = self.to_k_ip(ip_x)
ip_key = ip_key * torch.rsqrt(ip_key.pow(2).mean(dim=-1, keepdim=True) + 1e-5).to(ip_key.dtype)
ip_value = self.to_v_ip(ip_x)
ip_key = ip_key.view(b, -1, n, d)
ip_value = ip_value.view(b, -1, n, d)
ip_x = attention(q, ip_key, ip_value).reshape(b, -1, n * d)
return ip_x
class WanLynxRefAttention(nn.Module):
def __init__(self, dim=5120, bias=True):
super().__init__()
self.to_k_ref = nn.Linear(dim, dim, bias=bias)
self.to_v_ref = nn.Linear(dim, dim, bias=bias)
def forward(self, q, ref_feature: Optional[tuple] = None):
ref_x, ref_lens = ref_feature
ref_query = q
ref_query = self.self_attn.norm_q(ref_query)
ref_key = self.self_attn.norm_k(ref_key)
ref_query = ref_query.unflatten(2, (self.self_attn.heads, -1)).transpose(1, 2)
ref_key = ref_key.unflatten(2, (self.self_attn.heads, -1)).transpose(1, 2)
ref_value = ref_value.unflatten(2, (self.self_attn.heads, -1)).transpose(1, 2)
ref_x = attention(ref_query, ref_key, ref_value)
return self.self_attn.o(ref_x.flatten(2))
+219
View File
@@ -0,0 +1,219 @@
import os
import torch
import gc
from ..utils import log, dict_to_device
import numpy as np
from accelerate import init_empty_weights
from accelerate.utils import set_module_tensor_to_device
import comfy.model_management as mm
from comfy.utils import load_torch_file
import folder_paths
script_directory = os.path.dirname(os.path.abspath(__file__))
device = mm.get_torch_device()
offload_device = mm.unet_offload_device()
from .resampler import Resampler
class LoadLynxResampler:
@classmethod
def INPUT_TYPES(s):
return {
"required": {
"model_name": (folder_paths.get_filename_list("diffusion_models"), {"tooltip": "These models are loaded from 'ComfyUI/models/diffusion_models'"}),
"precision": (["fp32", "bf16", "fp16"], {"default": "fp16"}),
},
}
RETURN_TYPES = ("LYNXRESAMPLER",)
RETURN_NAMES = ("resampler", )
FUNCTION = "loadmodel"
CATEGORY = "WanVideoWrapper"
def loadmodel(self, model_name, precision):
dtype = {"bf16": torch.bfloat16, "fp16": torch.float16, "fp32": torch.float32}[precision]
model_path = folder_paths.get_full_path("diffusion_models", model_name)
resampler_sd = load_torch_file(model_path, safe_load=True)
output_dim = resampler_sd["proj_out.weight"].shape[0]
resampler = Resampler(
depth=4,
dim=1280,
dim_head=64,
embedding_dim=512,
ff_mult=4,
heads=20,
num_queries=16,
output_dim=output_dim,
dtype=dtype,
).eval()
resampler.to(offload_device, dtype)
resampler.load_state_dict(resampler_sd, strict=True)
#for name, param in resampler.named_parameters():
# print(f"{name}: {param.shape} {param.dtype}")
return resampler,
class VideoStyleInfo: # key names should match those used in style.yaml file
style_name: str = 'none'
num_frames: int = 81
seed: int = -1
guidance_scale: float = 5.0
guidance_scale_i: float = 2.0
num_inference_steps: int = 50
width: int = 832
height: int = 480
prompt: str = ''
negative_prompt: str = ''
class LynxEncodeFace:
@classmethod
def INPUT_TYPES(s):
return {
"required": {
"resampler": ("LYNXRESAMPLER", {"tooltip": "lynx resampler model"}),
"image": ("IMAGE", {"tooltip": "Input images for the model"}),
},
}
RETURN_TYPES = ("LYNXFACE", "IMAGE",)
RETURN_NAMES = ("lynx_face_embeds", "processed_image")
FUNCTION = "encode"
CATEGORY = "WanVideoWrapper"
def encode(self, resampler, image):
from .face.face_encoder import FaceEncoderArcFace, get_landmarks_from_image
image_np = (image[0].numpy() * 255).astype(np.uint8)
# Landmarks
landmarks = np.array([
[200.509, 213.98592],
[297.2495, 212.67685],
[245.74419, 272.85718],
[212.2043, 331.09564],
[288.75986, 330.27188]
])
# landmarks = np.array([[379.35895, 381.4803],
# [542.8676, 362.75436],
# [451.62717, 467.9116],
# [407.03708, 555.56305],
# [554.57874, 539.22833]])
#landmarks = get_landmarks_from_image(image_np)
print(landmarks)
# Face embedding via ArcFace
face_encoder = FaceEncoderArcFace()
face_encoder.init_encoder_model(device)
arcface_embed, processed_image = face_encoder(image_np, need_proc=True, landmarks=landmarks)
arcface_embed = arcface_embed.to(device, resampler.dtype)
arcface_embed = arcface_embed.reshape([1, -1, 512])
resampler.to(device)
ip_x = resampler(arcface_embed)
ip_x_uncond = resampler(arcface_embed * 0)
resampler.to(offload_device)
#ip_x = torch.load(os.path.join(script_directory, "debug_face_embeds.pt"))
#ip_x = ip_x[1].unsqueeze(0)
ip_x= ip_x.to(resampler.dtype)
out_dict = {
'ip_x': ip_x,
'ip_x_uncond': ip_x_uncond,
"landmarks": landmarks,
}
print("processed_image.shape", processed_image.min(), processed_image.max())
processed_image = (processed_image - processed_image.min()) / (processed_image.max() - processed_image.min())
processed_image = processed_image.permute(0, 2, 3, 1)
return out_dict, processed_image
class DrawArcFaceLandmarks:
@classmethod
def INPUT_TYPES(s):
return {
"required": {
"lynx_face_embeds": ("LYNXFACE", {"tooltip": "lynx resampler model"}),
"image": ("IMAGE", {"tooltip": "Input images for the model"}),
},
"optional": {
"image": ("IMAGE",)
}
}
RETURN_TYPES = ("IMAGE",)
RETURN_NAMES = ("landmarked_image", )
FUNCTION = "draw"
CATEGORY = "WanVideoWrapper"
def draw(self, lynx_face_embeds, image):
import cv2
landmarks = lynx_face_embeds['landmarks']
image_np = image[0].numpy() * 255
print("image_np.shape", image_np.shape) #image_np.shape (3, 512, 512)
print(type(landmarks))
print(landmarks)
for (x, y) in landmarks:
cv2.circle(image_np, (int(x), int(y)), radius=3, color=(0, 255, 0), thickness=-1)
image_out = torch.from_numpy(image_np / 255).unsqueeze(0).float()
print(image_out.shape) #torch.Size([1, 3, 512, 512])
return image_out,
class WanVideoAddLynxEmbeds:
@classmethod
def INPUT_TYPES(s):
return {"required": {
"embeds": ("WANVIDIMAGE_EMBEDS",),
"lynx_embeds": ("LYNXFACE", {"tooltip": "lynx face embeddings"}),
"strength": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 100.0, "step": 0.01, "tooltip": "Strength of the MTV motion"}),
"start_percent": ("FLOAT", {"default": 0.0, "min": 0.0, "max": 1.0, "step": 0.01, "tooltip": "Start percent to apply the ref "}),
"end_percent": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 1.0, "step": 0.01, "tooltip": "End percent to apply the ref "}),
}
}
RETURN_TYPES = ("WANVIDIMAGE_EMBEDS",)
RETURN_NAMES = ("image_embeds",)
FUNCTION = "add"
CATEGORY = "WanVideoWrapper"
def add(self, embeds, lynx_embeds, strength, start_percent, end_percent):
new_entry = {
"ip_x": lynx_embeds["ip_x"],
"ip_x_uncond": lynx_embeds["ip_x_uncond"],
"strength": strength,
"start_percent": start_percent,
"end_percent": end_percent,
}
updated = dict(embeds)
updated["lynx_embeds"] = new_entry
return (updated,)
NODE_CLASS_MAPPINGS = {
"LoadLynxResampler": LoadLynxResampler,
"LynxEncodeFace": LynxEncodeFace,
"DrawArcFaceLandmarks": DrawArcFaceLandmarks,
"WanVideoAddLynxEmbeds": WanVideoAddLynxEmbeds,
}
NODE_DISPLAY_NAME_MAPPINGS = {
"LoadLynxResampler": "Load Lynx Resampler",
"LynxEncodeFace": "Lynx Encode Face",
"DrawArcFaceLandmarks": "Draw ArcFace Landmarks",
"WanVideoAddLynxEmbeds": "WanVideo Add Lynx Embeds",
}
+154
View File
@@ -0,0 +1,154 @@
# Copyright (c) 2022 Phil Wang
# Copyright (c) 2023 Anas Awadalla, Irena Gao, Joshua Gardner, Jack Hessel, Yusuf Hanafy, Wanrong Zhu, Kalyani Marathe, Yonatan Bitton, Samir Gadre, Jenia Jitsev, Simon Kornblith, Pang Wei Koh, Gabriel Ilharco, Mitchell Wortsman, Ludwig Schmidt.
# Copyright (c) 2023 Tencent AI Lab
# Copyright (c) 2025 Bytedance Ltd. and/or its affiliates
# This file has been modified by Bytedance Ltd. and/or its affiliates on September 15, 2025.
# SPDX-License-Identifier: Apache-2.0
# Original file (IP-Adapter) was released under Apache License 2.0, with the full license text
# available at https://github.com/tencent-ailab/IP-Adapter/blob/main/LICENSE.
# Original file (open_flamingo and flamingo-pytorch) was released under MIT License:
# Permission is hereby granted, free of charge, to any person obtaining a copy
# of this software and associated documentation files (the "Software"), to deal
# in the Software without restriction, including without limitation the rights
# to use, copy, modify, merge, publish, distribute, sublicense, and/or sell
# copies of the Software, and to permit persons to whom the Software is
# furnished to do so, subject to the following conditions:
# The above copyright notice and this permission notice shall be included in all
# copies or substantial portions of the Software.
# THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR
# IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY,
# FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE
# AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER
# LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM,
# OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN THE
# SOFTWARE.
import math
import torch
import torch.nn as nn
# FFN
def FeedForward(dim, mult=4):
inner_dim = int(dim * mult)
return nn.Sequential(
nn.LayerNorm(dim),
nn.Linear(dim, inner_dim, bias=False),
nn.GELU(),
nn.Linear(inner_dim, dim, bias=False),
)
def reshape_tensor(x, heads):
bs, length, width = x.shape
#(bs, length, width) --> (bs, length, n_heads, dim_per_head)
x = x.view(bs, length, heads, -1)
# (bs, length, n_heads, dim_per_head) --> (bs, n_heads, length, dim_per_head)
x = x.transpose(1, 2)
# (bs, n_heads, length, dim_per_head) --> (bs*n_heads, length, dim_per_head)
x = x.reshape(bs, heads, length, -1)
return x
class PerceiverAttention(nn.Module):
def __init__(self, *, dim, dim_head=64, heads=8):
super().__init__()
self.scale = dim_head**-0.5
self.dim_head = dim_head
self.heads = heads
inner_dim = dim_head * heads
self.norm1 = nn.LayerNorm(dim)
self.norm2 = nn.LayerNorm(dim)
self.to_q = nn.Linear(dim, inner_dim, bias=False)
self.to_kv = nn.Linear(dim, inner_dim * 2, bias=False)
self.to_out = nn.Linear(inner_dim, dim, bias=False)
def forward(self, x, latents):
"""
Args:
x (torch.Tensor): image features
shape (b, n1, D)
latent (torch.Tensor): latent features
shape (b, n2, D)
"""
x = self.norm1(x)
latents = self.norm2(latents)
b, l, _ = latents.shape
q = self.to_q(latents)
kv_input = torch.cat((x, latents), dim=-2)
k, v = self.to_kv(kv_input).chunk(2, dim=-1)
q = reshape_tensor(q, self.heads)
k = reshape_tensor(k, self.heads)
v = reshape_tensor(v, self.heads)
# attention
scale = 1 / math.sqrt(math.sqrt(self.dim_head))
weight = (q * scale) @ (k * scale).transpose(-2, -1) # More stable with f16 than dividing afterwards
weight = torch.softmax(weight.float(), dim=-1).type(weight.dtype)
out = weight @ v
out = out.permute(0, 2, 1, 3).reshape(b, l, -1)
return self.to_out(out)
class Resampler(nn.Module):
def __init__(
self,
dim=1024,
depth=8,
dim_head=64,
heads=16,
num_queries=8,
embedding_dim=768,
output_dim=1024,
ff_mult=4,
dtype=torch.float32,
):
super().__init__()
self.latents = nn.Parameter(torch.randn(1, num_queries, dim) / dim**0.5)
self.proj_in = nn.Linear(embedding_dim, dim)
self.proj_out = nn.Linear(dim, output_dim)
self.norm_out = nn.LayerNorm(output_dim)
self.dtype = dtype
self.layers = nn.ModuleList([])
for _ in range(depth):
self.layers.append(
nn.ModuleList(
[
PerceiverAttention(dim=dim, dim_head=dim_head, heads=heads),
FeedForward(dim=dim, mult=ff_mult),
]
)
)
def forward(self, x):
latents = self.latents.repeat(x.size(0), 1, 1)
x = self.proj_in(x)
for attn, ff in self.layers:
latents = attn(x, latents) + latents
latents = ff(latents) + latents
latents = self.proj_out(latents)
return self.norm_out(latents)
+32 -18
View File
@@ -1141,6 +1141,15 @@ class WanVideoModelLoader:
is_humo = "audio_proj.audio_proj_glob_1.layer.weight" in sd
is_wananimate = "pose_patch_embedding.weight" in sd
#lynx
lynx_layers = "none"
if "blocks.0.cross_attn.ip_adapter.to_v_ip.weight" in sd and "blocks.0.ref_adapter.to_k_ref.weight" in sd:
n_registers = sd["blocks.0.cross_attn.ip_adapter.registers"].shape[1]
lynx_layers = "full"
elif "blocks.0.cross_attn.ip_adapter.to_v_ip.weight" in sd:
n_registers = 0
lynx_layers = "lite"
model_type = "t2v"
if "audio_injector.injector.0.k.weight" in sd:
model_type = "s2v"
@@ -1259,6 +1268,7 @@ class WanVideoModelLoader:
"humo_audio": is_humo,
"is_wananimate": is_wananimate,
"rms_norm_function": rms_norm_function,
"lynx_layers": lynx_layers,
}
@@ -1399,27 +1409,31 @@ class WanVideoModelLoader:
del unianimate_sd
if not gguf:
if merge_loras and lora is not None:
if not lora_low_mem_load:
load_weights(transformer, sd, weight_dtype, base_dtype, transformer_load_device)
if control_lora:
patch_control_lora(patcher.model.diffusion_model, device)
patcher.model.is_patched = True
if lora is not None:
if merge_loras:
if not lora_low_mem_load:
load_weights(transformer, sd, weight_dtype, base_dtype, transformer_load_device)
log.info("Merging LoRA to the model...")
patcher = apply_lora(
patcher, device, transformer_load_device, params_to_keep=params_to_keep, dtype=weight_dtype, base_dtype=base_dtype, state_dict=sd,
low_mem_load=lora_low_mem_load, control_lora=control_lora, scale_weights=scale_weights)
if not control_lora:
scale_weights.clear()
patcher.patches.clear()
if control_lora:
patch_control_lora(patcher.model.diffusion_model, device)
patcher.model.is_patched = True
log.info("Merging LoRA to the model...")
patcher = apply_lora(
patcher, device, transformer_load_device, params_to_keep=params_to_keep, dtype=weight_dtype, base_dtype=base_dtype, state_dict=sd,
low_mem_load=lora_low_mem_load, control_lora=control_lora, scale_weights=scale_weights)
if not control_lora:
scale_weights.clear()
patcher.patches.clear()
transformer.patched_linear = False
sd = None
else:
from .custom_linear import _replace_linear
transformer = _replace_linear(transformer, base_dtype, sd, scale_weights=scale_weights)
else:
load_weights(transformer, sd, weight_dtype, base_dtype, transformer_load_device)
transformer.patched_linear = False
sd = None
else:
from .custom_linear import _replace_linear
transformer = _replace_linear(transformer, base_dtype, sd, scale_weights=scale_weights)
transformer.patched_linear = True
if "fast" in quantization:
if lora is not None and not merge_loras:
+7 -1
View File
@@ -615,6 +615,11 @@ class WanVideoSampler:
"end_percent": fantasy_portrait_embeds.get("end_percent", 1.0),
}
# Lynx
lynx_embeds = image_embeds.get("lynx_embeds", None)
if lynx_embeds is not None:
log.info("Using Lynx embeddings", lynx_embeds)
# MiniMax Remover
minimax_latents = minimax_mask_latents = None
minimax_latents = image_embeds.get("minimax_latents", None)
@@ -1229,7 +1234,8 @@ class WanVideoSampler:
"wananim_pose_latents": wananim_pose_latents.to(device) if wananim_pose_latents is not None else None, # WanAnimate pose latents
"wananim_face_pixel_values": wananim_face_pixels.to(device, torch.float32) if wananim_face_pixels is not None else None, # WanAnimate face images
"wananim_pose_strength": wananim_pose_strength,
"wananim_face_strength": wananim_face_strength
"wananim_face_strength": wananim_face_strength,
"lynx_embeds": lynx_embeds, # Lynx face and reference embeddings
}
batch_size = 1
+39 -6
View File
@@ -620,11 +620,12 @@ class WanT2VCrossAttention(WanSelfAttention):
def __init__(self, in_features, out_features, num_heads, kv_dim=None, qk_norm=True, eps=1e-6, attention_mode='sdpa', rms_norm_function="default"):
super().__init__(in_features, out_features, num_heads, qk_norm, eps, kv_dim=kv_dim, rms_norm_function=rms_norm_function)
self.attention_mode = attention_mode
self.ip_adapter = None
def forward(self, x, context, grid_sizes=None, clip_embed=None, audio_proj=None, audio_scale=1.0,
num_latent_frames=21, nag_params={}, nag_context=None, is_uncond=False, rope_func="comfy",
inner_t=None, inner_c=None, cross_freqs=None,
adapter_proj=None, adapter_attn_mask=None, ip_scale=1.0, orig_seq_len=None, **kwargs):
adapter_proj=None, adapter_attn_mask=None, ip_scale=1.0, orig_seq_len=None, lynx_x_ip=None, lynx_ip_scale=1.0, **kwargs):
b, n, d = x.size(0), self.num_heads, self.head_dim
# compute query
q = self.norm_q(self.q(x),num_chunks=2 if rope_func == "comfy_chunked" else 1).view(b, -1, n, d)
@@ -645,6 +646,10 @@ class WanT2VCrossAttention(WanSelfAttention):
x = x_text
if lynx_x_ip is not None and self.ip_adapter is not None and ip_scale !=0:
lynx_x_ip = self.ip_adapter(self, q, x, lynx_x_ip)
x = x.add(lynx_x_ip, alpha=lynx_ip_scale)
# FantasyTalking audio attention
if audio_proj is not None:
if len(audio_proj.shape) == 4:
@@ -840,7 +845,9 @@ class WanAttentionBlock(nn.Module):
rms_norm_function="default",
use_motion_attn=False,
use_humo_audio_attn=False,
face_fuser_block=False
face_fuser_block=False,
lynx_layers="none",
block_idx=0
):
super().__init__()
self.dim = out_features
@@ -855,6 +862,7 @@ class WanAttentionBlock(nn.Module):
self.dense_timesteps = 10
self.dense_block = False
self.dense_attention_mode = "sageattn"
self.block_idx = block_idx
self.kv_cache = None
self.use_motion_attn = use_motion_attn
@@ -889,6 +897,17 @@ class WanAttentionBlock(nn.Module):
from .wananimate.face_blocks import FaceBlock
self.fuser_block = FaceBlock(self.dim, num_heads)
# Lynx
self.ref_adapter = None
if lynx_layers == "full":
from ...lynx.modules import WanLynxIPCrossAttention, WanLynxRefAttention
self.cross_attn.ip_adapter = WanLynxIPCrossAttention(cross_attention_dim=self.dim, dim=self.dim, n_registers=16)
self.ref_adapter = WanLynxRefAttention(dim=self.dim)
elif lynx_layers == "lite":
from ...lynx.modules import WanLynxIPCrossAttention
if self.block_idx % 2 == 0:
self.cross_attn.ip_adapter = WanLynxIPCrossAttention(cross_attention_dim=2048, dim=self.dim, n_registers=0, bias=False)
#@torch.compiler.disable()
def get_mod(self, e):
if e.dim() == 3:
@@ -968,6 +987,7 @@ class WanAttentionBlock(nn.Module):
reverse_time=False,
mtv_motion_tokens=None, mtv_motion_rotary_emb=None, mtv_strength=1.0, mtv_freqs=None,
humo_audio_input=None, humo_audio_scale=1.0,
lynx_x_ip=None, lynx_x_ref=None, lynx_ip_scale = 1.0, lynx_ref_scale=1.0,
):
r"""
Args:
@@ -1126,7 +1146,7 @@ class WanAttentionBlock(nn.Module):
multitalk_audio_embedding, x_ref_attn_map, human_num, inner_t, inner_c, cross_freqs,
adapter_proj=adapter_proj, ip_scale=ip_scale,
mtv_freqs=mtv_freqs, mtv_motion_tokens=mtv_motion_tokens, mtv_motion_rotary_emb=mtv_motion_rotary_emb, mtv_strength=mtv_strength,
humo_audio_input=humo_audio_input, humo_audio_scale=humo_audio_scale
humo_audio_input=humo_audio_input, humo_audio_scale=humo_audio_scale, lynx_x_ip=lynx_x_ip, lynx_ip_scale=lynx_ip_scale
)
else:
if self.rope_func == "comfy_chunked":
@@ -1148,13 +1168,13 @@ class WanAttentionBlock(nn.Module):
audio_proj, audio_scale, num_latent_frames, nag_params,
nag_context, is_uncond, multitalk_audio_embedding, x_ref_attn_map, human_num,
inner_t, inner_c, cross_freqs, adapter_proj, ip_scale, mtv_freqs, mtv_motion_tokens, mtv_motion_rotary_emb, mtv_strength,
humo_audio_input, humo_audio_scale):
humo_audio_input, humo_audio_scale, lynx_x_ip, lynx_ip_scale):
x = x + self.cross_attn(self.norm3(x), context, grid_sizes, clip_embed=clip_embed,
audio_proj=audio_proj, audio_scale=audio_scale,
num_latent_frames=num_latent_frames, nag_params=nag_params, nag_context=nag_context, is_uncond=is_uncond,
rope_func=self.rope_func, inner_t=inner_t, inner_c=inner_c, cross_freqs=cross_freqs,
adapter_proj=adapter_proj, ip_scale=ip_scale, orig_seq_len=self.original_seq_len)
adapter_proj=adapter_proj, ip_scale=ip_scale, orig_seq_len=self.original_seq_len, lynx_x_ip=lynx_x_ip, lynx_ip_scale=lynx_ip_scale)
# MultiTalk
if multitalk_audio_embedding is not None and not isinstance(self, VaceWanAttentionBlock):
x_audio = self.audio_cross_attn(self.norm_x(x), encoder_hidden_states=multitalk_audio_embedding,
@@ -1502,6 +1522,8 @@ class WanModel(torch.nn.Module):
# WanAnimate
is_wananimate=False,
motion_encoder_dim=512,
# lynx
lynx_layers="none",
):
r"""
Initialize the diffusion model backbone.
@@ -1669,7 +1691,7 @@ class WanModel(torch.nn.Module):
qk_norm, cross_attn_norm, eps,
attention_mode=self.attention_mode, rope_func=self.rope_func, rms_norm_function=rms_norm_function,
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))
face_fuser_block = (i % 5 == 0 and is_wananimate), lynx_layers=lynx_layers, block_idx=i)
for i in range(num_layers)
])
#MTV Crafter
@@ -2044,6 +2066,7 @@ class WanModel(torch.nn.Module):
wananim_face_pixel_values=None,
wananim_pose_strength=1.0,
wananim_face_strength=1.0,
lynx_embeds=None,
):
r"""
@@ -2089,6 +2112,14 @@ class WanModel(torch.nn.Module):
if hasattr(submodule, 'step'):
submodule.step = current_step
lynx_x_ip = None
if lynx_embeds is not None:
if not is_uncond:
lynx_x_ip = lynx_embeds["ip_x"].to(self.main_device)
else:
lynx_x_ip = lynx_embeds["ip_x_uncond"].to(self.main_device)
lynx_ip_scale = lynx_embeds.get("strength", 1.0)
#s2v
if self.model_type == 's2v' and s2v_audio_input is not None:
if is_uncond:
@@ -2574,6 +2605,8 @@ class WanModel(torch.nn.Module):
mtv_motion_tokens=mtv_motion_tokens, mtv_motion_rotary_emb=mtv_motion_rotary_emb, mtv_strength=mtv_strength, mtv_freqs=mtv_freqs,
humo_audio_input=humo_audio_input,
humo_audio_scale=humo_audio_scale,
lynx_x_ip=lynx_x_ip,
lynx_ip_scale=lynx_ip_scale,
)
if vace_data is not None: