init: support lite model
This commit is contained in:
@@ -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"]
|
||||
@@ -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
|
||||
@@ -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
|
||||
@@ -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
@@ -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",
|
||||
}
|
||||
@@ -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)
|
||||
+16
-2
@@ -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,7 +1409,8 @@ class WanVideoModelLoader:
|
||||
del unianimate_sd
|
||||
|
||||
if not gguf:
|
||||
if merge_loras and lora is not None:
|
||||
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)
|
||||
|
||||
@@ -1419,7 +1430,10 @@ class WanVideoModelLoader:
|
||||
else:
|
||||
from .custom_linear import _replace_linear
|
||||
transformer = _replace_linear(transformer, base_dtype, sd, scale_weights=scale_weights)
|
||||
transformer.patched_linear = True
|
||||
else:
|
||||
load_weights(transformer, sd, weight_dtype, base_dtype, transformer_load_device)
|
||||
transformer.patched_linear = False
|
||||
sd = None
|
||||
|
||||
if "fast" in quantization:
|
||||
if lora is not None and not merge_loras:
|
||||
|
||||
+7
-1
@@ -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
|
||||
|
||||
@@ -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:
|
||||
|
||||
Reference in New Issue
Block a user