7 Commits
Author SHA1 Message Date
kijai 7a0da7708e Merge branch 'main' into vap 2025-11-04 23:15:20 +02:00
kijai a51a53d5b7 Use proper self_attn output layers, and cleanup 2025-11-04 21:13:29 +02:00
kijai 977f4a5c3a Update model.py 2025-11-04 11:31:08 +02:00
kijai 2a45675498 Merge branch 'main' into vap 2025-11-04 10:39:35 +02:00
kijai ea414c54ac Create wanvideo_I2V_video-as-prompt_testing_WIP.json 2025-11-01 16:53:15 +02:00
kijai 0013ae0ece Init VAP 2025-11-01 16:47:37 +02:00
kijai 0e904e6035 Remove unnecessary casts 2025-10-31 17:28:12 +02:00
10 changed files with 2188 additions and 3415 deletions
+3 -13
View File
@@ -8,8 +8,6 @@ from .vae.autoencoder import AutoEncoderModule
from .vae.distributions import DiagonalGaussianDistribution
import torchaudio
from ..utils import log
from comfy import model_management as mm
device = mm.get_torch_device()
offload_device = mm.unet_offload_device()
@@ -218,11 +216,9 @@ class WanVideoOviCFG:
def INPUT_TYPES(s):
return {"required": {
"original_text_embeds": ("WANVIDEOTEXTEMBEDS",),
"ovi_negative_text_embeds": ("WANVIDEOTEXTEMBEDS",),
"ovi_audio_cfg": ("FLOAT", {"default": 3.0, "min": 0.0, "max": 100.0, "step": 0.01}),
},
"optional": {
"ovi_negative_text_embeds": ("WANVIDEOTEXTEMBEDS",),
}
}
RETURN_TYPES = ("WANVIDEOTEXTEMBEDS", )
@@ -231,16 +227,10 @@ class WanVideoOviCFG:
CATEGORY = "WanVideoWrapper/Ovi"
DESCRIPTION = "Adds Ovi negative text embeddings and audio CFG scale to the text embeddings dictionary"
def process(self, original_text_embeds, ovi_audio_cfg, ovi_negative_text_embeds=None):
negative_text_embeds = None
if ovi_negative_text_embeds is not None:
negative_text_embeds = ovi_negative_text_embeds.get("prompt_embeds", None)
def process(self, original_text_embeds, ovi_negative_text_embeds, ovi_audio_cfg):
negative_text_embeds = ovi_negative_text_embeds.get("negative_prompt_embeds", None)
if negative_text_embeds is None:
negative_text_embeds = original_text_embeds["prompt_embeds"]
log.info("WanVideoOviCFG: Ovi negative text embeddings not provided, using original prompt embeddings as negative embeddings")
else:
log.info("WanVideoOviCFG: Using provided Ovi audio negative text embeddings")
log.info("WanVideoOviCFG: negative text embedding shape: {}".format(negative_text_embeds[0].shape))
prompt_embeds_dict_copy = original_text_embeds.copy()
prompt_embeds_dict_copy.update({
+31 -139
View File
@@ -22,6 +22,30 @@ offload_device = mm.unet_offload_device()
VAE_STRIDE = (4, 8, 8)
PATCH_SIZE = (1, 2, 2)
class WanVideoAddVideoPromptEmbeds:
@classmethod
def INPUT_TYPES(s):
return {"required": {
"image_embeds": ("WANVIDIMAGE_EMBEDS",),
"video_prompt_embeds": ("WANVIDIMAGE_EMBEDS",),
"video_prompt_latents": ("LATENT", ),
"text_embeds": ("WANVIDEOTEXTEMBEDS", ),
}
}
RETURN_TYPES = ("WANVIDIMAGE_EMBEDS",)
RETURN_NAMES = ("image_embeds",)
FUNCTION = "add"
CATEGORY = "WanVideoWrapper"
EXPERIMENTAL = True
def add(self, image_embeds, video_prompt_embeds, video_prompt_latents, text_embeds):
updated = dict(image_embeds)
updated["video_prompt_embeds"] = video_prompt_embeds
updated["video_prompt_embeds"]["video_prompt_latents"] = video_prompt_latents["samples"][0]
updated["video_prompt_embeds"]["text_embeds"] = text_embeds
return (updated,)
class WanVideoEnhanceAVideo:
@classmethod
@@ -765,98 +789,6 @@ class WanVideoAddStandInLatent:
updated = dict(embeds)
updated["standin_input"] = new_entry
return (updated,)
class WanVideoAddBindweaveEmbeds:
@classmethod
def INPUT_TYPES(s):
return {"required": {
"embeds": ("WANVIDIMAGE_EMBEDS",),
"reference_latents": ("LATENT", {"tooltip": "Reference image to encode"}),
},
"optional": {
"ref_masks": ("MASK", {"tooltip": "Reference mask to encode"}),
"qwenvl_embeds_pos": ("QWENVL_EMBEDS", {"tooltip": "Qwen-VL image embeddings for the reference image"}),
"qwenvl_embeds_neg": ("QWENVL_EMBEDS", {"tooltip": "Qwen-VL image embeddings for the reference image"}),
}
}
RETURN_TYPES = ("WANVIDIMAGE_EMBEDS", "LATENT", "MASK",)
RETURN_NAMES = ("image_embeds", "image_embed_preview", "mask_preview",)
FUNCTION = "add"
CATEGORY = "WanVideoWrapper"
def add(self, embeds, reference_latents, ref_masks=None, qwenvl_embeds_pos=None, qwenvl_embeds_neg=None):
updated = dict(embeds)
image_embeds = embeds["image_embeds"]
max_refs = 4
num_refs = reference_latents["samples"].shape[0]
pad = torch.zeros(image_embeds.shape[0], max_refs-num_refs, image_embeds.shape[2], image_embeds.shape[3], device=image_embeds.device, dtype=image_embeds.dtype)
if num_refs < max_refs:
image_embeds = torch.cat([pad, image_embeds], dim=1)
ref_latents = [ref_latent for ref_latent in reference_latents["samples"]]
image_embeds = torch.cat([*ref_latents, image_embeds], dim=1)
mask = embeds.get("mask", None)
if mask is not None:
mask_pad = torch.zeros(mask.shape[0], max_refs-num_refs, mask.shape[2], mask.shape[3], device=mask.device, dtype=mask.dtype)
if num_refs < max_refs:
mask = torch.cat([mask_pad, mask], dim=1)
if ref_masks is not None:
ref_mask_ = common_upscale(ref_masks.unsqueeze(1), mask.shape[3], mask.shape[2], "nearest", "disabled").movedim(0,1)
ref_mask_ = torch.cat([ref_mask_, torch.zeros(3, ref_mask_.shape[1], ref_mask_.shape[2], ref_mask_.shape[3], device=ref_mask_.device, dtype=ref_mask_.dtype)])
mask = torch.cat([ref_mask_, mask], dim=1)
else:
mask = torch.cat([torch.ones(mask.shape[0], num_refs, mask.shape[2], mask.shape[3], device=mask.device, dtype=mask.dtype), mask], dim=1)
updated["mask"] = mask
clip_embeds = updated.get("clip_context", None)
if clip_embeds is not None:
B, T, C = clip_embeds.shape
target_len = max_refs * 257 # 4 * 257 = 1028
if T < target_len:
pad = torch.zeros(B, target_len - T, C, device=clip_embeds.device, dtype=clip_embeds.dtype)
padded_embeds = torch.cat([clip_embeds, pad], dim=1)
log.info(f"Padded clip embeds from {clip_embeds.shape} to {padded_embeds.shape} for Bindweave")
updated["clip_context"] = padded_embeds
else:
updated["clip_context"] = clip_embeds
updated["image_embeds"] = image_embeds
updated["qwenvl_embeds_pos"] = qwenvl_embeds_pos
updated["qwenvl_embeds_neg"] = qwenvl_embeds_neg
return (updated, {"samples": image_embeds.unsqueeze(0)}, mask[0].float())
class TextImageEncodeQwenVL():
@classmethod
def INPUT_TYPES(s):
return {"required": {
"clip": ("CLIP",),
"prompt": ("STRING", {"default": "", "multiline": True}),
},
"optional": {
"image": ("IMAGE", ),
}
}
RETURN_TYPES = ("QWENVL_EMBEDS",)
RETURN_NAMES = ("qwenvl_embeds",)
FUNCTION = "add"
CATEGORY = "WanVideoWrapper"
def add(cls, clip, prompt, image=None):
if image is None:
input_images = []
llama_template = None
else:
input_images = [image[:, :, :, :3]]
llama_template = "<|im_start|>system\nYou are a helpful assistant.<|im_end|>\n<|im_start|>user\n<|vision_start|><|image_pad|><|vision_end|>{}<|im_end|>\n<|im_start|>assistant\n"
tokens = clip.tokenize(prompt, images=input_images, llama_template=llama_template)
conditioning = clip.encode_from_tokens_scheduled(tokens)
print("Qwen-VL embeds shape:", conditioning[0][0].shape)
return (conditioning[0][0],)
class WanVideoAddMTVMotion:
@classmethod
@@ -927,6 +859,10 @@ class WanVideoImageToVideoEncode:
start_latent_strength, end_latent_strength, start_image=None, end_image=None, control_embeds=None, fun_or_fl2v_model=False,
temporal_mask=None, extra_latents=None, clip_embeds=None, tiled_vae=False, add_cond_latents=None, vae=None):
if start_image is None and end_image is None and add_cond_latents is None:
return WanVideoEmptyEmbeds().process(
num_frames, width, height, control_embeds=control_embeds, extra_latents=extra_latents,
)
if vae is None:
raise ValueError("VAE is required for image encoding.")
H = height
@@ -1044,7 +980,7 @@ class WanVideoImageToVideoEncode:
gc.collect()
image_embeds = {
"image_embeds": y.cpu(),
"image_embeds": y,
"clip_context": clip_embeds.get("clip_embeds", None) if clip_embeds is not None else None,
"negative_clip_context": clip_embeds.get("negative_clip_embeds", None) if clip_embeds is not None else None,
"max_seq_len": max_seq_len,
@@ -1056,7 +992,7 @@ class WanVideoImageToVideoEncode:
"fun_or_fl2v_model": fun_or_fl2v_model,
"has_ref": has_ref,
"add_cond_latents": add_cond_latents,
"mask": mask.cpu()
"mask": mask
}
return (image_embeds,)
@@ -1246,46 +1182,6 @@ class WanVideoAnimateEmbeds:
}
return (image_embeds,)
# region UniLumos
class WanVideoUniLumosEmbeds:
@classmethod
def INPUT_TYPES(s):
return {"required": {
"width": ("INT", {"default": 832, "min": 64, "max": 8096, "step": 8, "tooltip": "Width of the image to encode"}),
"height": ("INT", {"default": 480, "min": 64, "max": 8096, "step": 8, "tooltip": "Height of the image to encode"}),
"num_frames": ("INT", {"default": 81, "min": 1, "max": 10000, "step": 4, "tooltip": "Number of frames to encode"}),
},
"optional": {
"foreground_latents": ("LATENT", {"tooltip": "Video foreground latents"}),
"background_latents": ("LATENT", {"tooltip": "Video background latents"}),
}
}
RETURN_TYPES = ("WANVIDIMAGE_EMBEDS", )
RETURN_NAMES = ("image_embeds",)
FUNCTION = "process"
CATEGORY = "WanVideoWrapper"
def process(self, num_frames, width, height, foreground_latents=None, background_latents=None):
target_shape = (16, (num_frames - 1) // VAE_STRIDE[0] + 1,
height // VAE_STRIDE[1],
width // VAE_STRIDE[2])
embeds = {
"target_shape": target_shape,
"num_frames": num_frames,
}
if foreground_latents is not None:
embeds["foreground_latents"] = foreground_latents["samples"][0]
else:
embeds["foreground_latents"] = torch.zeros(target_shape[0], target_shape[1], target_shape[2], target_shape[3], device=torch.device("cpu"), dtype=torch.float32)
if background_latents is not None:
embeds["background_latents"] = background_latents["samples"][0]
else:
embeds["background_latents"] = torch.zeros(target_shape[0], target_shape[1], target_shape[2], target_shape[3], device=torch.device("cpu"), dtype=torch.float32)
return (embeds,)
class WanVideoEmptyEmbeds:
@classmethod
@@ -2334,9 +2230,7 @@ NODE_CLASS_MAPPINGS = {
"WanVideoAnimateEmbeds": WanVideoAnimateEmbeds,
"WanVideoAddLucyEditLatents": WanVideoAddLucyEditLatents,
"WanVideoSchedulerSA_ODE": WanVideoSchedulerSA_ODE,
"WanVideoAddBindweaveEmbeds": WanVideoAddBindweaveEmbeds,
"TextImageEncodeQwenVL": TextImageEncodeQwenVL,
"WanVideoUniLumosEmbeds": WanVideoUniLumosEmbeds,
"WanVideoAddVideoPromptEmbeds": WanVideoAddVideoPromptEmbeds,
}
NODE_DISPLAY_NAME_MAPPINGS = {
@@ -2376,6 +2270,4 @@ NODE_DISPLAY_NAME_MAPPINGS = {
"WanVideoAnimateEmbeds": "WanVideo Animate Embeds",
"WanVideoAddLucyEditLatents": "WanVideo Add LucyEdit Latents",
"WanVideoSchedulerSA_ODE": "WanVideo Scheduler SA-ODE",
"WanVideoAddBindweaveEmbeds": "WanVideo Add Bindweave Embeds",
"WanVideoUniLumosEmbeds": "WanVideo UniLumos Embeds",
}
+11 -23
View File
@@ -395,7 +395,7 @@ class WanVideoLoraSelect:
return (loras_list,)
try:
lora_path = folder_paths.get_full_path_or_raise("loras", lora)
lora_path = folder_paths.get_full_path("loras", lora)
except:
lora_path = lora
@@ -532,7 +532,7 @@ class WanVideoLoraSelectMulti:
if not lora_name or lora_name == "none" or s == 0.0:
continue
loras_list.append({
"path": folder_paths.get_full_path_or_raise("loras", lora_name),
"path": folder_paths.get_full_path("loras", lora_name),
"strength": s,
"name": os.path.splitext(lora_name)[0],
"blocks": blocks.get("selected_blocks", {}),
@@ -560,7 +560,7 @@ class WanVideoVACEModelSelect:
DESCRIPTION = "VACE model to use when not using model that has it included, loaded from 'ComfyUI/models/diffusion_models'"
def getvacepath(self, vace_model):
vace_model = [{"path": folder_paths.get_full_path_or_raise("diffusion_models", vace_model)}]
vace_model = [{"path": folder_paths.get_full_path("diffusion_models", vace_model)}]
return (vace_model,)
class WanVideoExtraModelSelect:
@@ -582,7 +582,7 @@ class WanVideoExtraModelSelect:
DESCRIPTION = "Extra model to load and add to the main model, ie. VACE or MTV Crafter 'ComfyUI/models/diffusion_models'"
def getmodelpath(self, extra_model, prev_model=None):
extra_model = {"path": folder_paths.get_full_path_or_raise("diffusion_models", extra_model)}
extra_model = {"path": folder_paths.get_full_path("diffusion_models", extra_model)}
if prev_model is not None and isinstance(prev_model, list):
extra_model_list = prev_model + [extra_model]
else:
@@ -1088,14 +1088,6 @@ class WanVideoModelLoader:
sd, reader = load_gguf(model_path)
gguf_reader.append(reader)
# Ovi
extra_audio_model = False
if any(key.startswith("video_model.") for key in sd.keys()):
sd = {key.replace("video_model.", "", 1).replace("modulation.modulation", "modulation"): value for key, value in sd.items()}
if any(key.startswith("audio_model.") for key in sd.keys()) and any(key.startswith("blocks.") for key in sd.keys()):
extra_audio_model = True
is_wananimate = "pose_patch_embedding.weight" in sd
# rename WanAnimate face fuser block keys to insert into main blocks instead
if is_wananimate:
@@ -1148,6 +1140,7 @@ class WanVideoModelLoader:
raise ValueError("You are attempting to load a VACE module as a WanVideo model, instead you should use the vace_model input and matching T2V base model")
# currently this can be VACE, MTV-Crafter, Lynx or Ovi-audio weights
extra_audio_model = False
if extra_model is not None:
for _model in extra_model:
print("Loading extra model: ", _model["path"])
@@ -1356,6 +1349,7 @@ class WanVideoModelLoader:
"lynx_ip_layers": lynx_ip_layers,
"lynx_ref_layers": lynx_ref_layers,
"is_longcat": dim == 4096,
"is_VAP": True if "patch_embedding_mot_ref.weight" in sd else False
}
@@ -1486,12 +1480,6 @@ class WanVideoModelLoader:
transformer.add_proj = zero_module(torch.nn.Linear(inner_dim, inner_dim))
transformer.attn_conv_in = torch.nn.Conv3d(attn_cond_in_dim, inner_dim, kernel_size=transformer.patch_size, stride=transformer.patch_size)
# Bindweave text_projection
if "text_projection.0.weight" in sd:
log.info("Bindweave model detected, adding text_projection to the model")
text_dim = sd["text_projection.0.weight"].shape[0]
transformer.text_projection = nn.Sequential(nn.Linear(sd["text_projection.0.weight"].shape[1], text_dim), nn.GELU(approximate='tanh'), nn.Linear(text_dim, text_dim))
latent_format=Wan22 if dim == 3072 else Wan21
comfy_model = WanVideoModel(
WanVideoModelConfig(base_dtype, latent_format=latent_format),
@@ -1689,7 +1677,7 @@ class WanVideoVAELoader:
def loadmodel(self, model_name, precision, compile_args=None):
dtype = {"bf16": torch.bfloat16, "fp16": torch.float16, "fp32": torch.float32}[precision]
model_path = folder_paths.get_full_path_or_raise("vae", model_name)
model_path = folder_paths.get_full_path("vae", model_name)
vae_sd = load_torch_file(model_path, safe_load=True)
has_model_prefix = any(k.startswith("model.") for k in vae_sd.keys())
@@ -1740,7 +1728,7 @@ class WanVideoTinyVAELoader:
from .taehv import TAEHV
dtype = {"bf16": torch.bfloat16, "fp16": torch.float16, "fp32": torch.float32}[precision]
model_path = folder_paths.get_full_path_or_raise("vae_approx", model_name)
model_path = folder_paths.get_full_path("vae_approx", model_name)
vae_sd = load_torch_file(model_path, safe_load=True)
vae = TAEHV(vae_sd, parallel=parallel, dtype=dtype)
@@ -1778,7 +1766,7 @@ class LoadWanVideoT5TextEncoder:
dtype = {"bf16": torch.bfloat16, "fp16": torch.float16, "fp32": torch.float32}[precision]
model_path = folder_paths.get_full_path_or_raise("text_encoders", model_name)
model_path = folder_paths.get_full_path("text_encoders", model_name)
sd = load_torch_file(model_path, safe_load=True)
if quantization == "disabled":
@@ -1888,10 +1876,10 @@ class LoadWanVideoClipTextEncoder:
dtype = {"bf16": torch.bfloat16, "fp16": torch.float16, "fp32": torch.float32}[precision]
model_path = folder_paths.get_full_path_or_raise("clip_vision", model_name)
model_path = folder_paths.get_full_path("clip_vision", model_name)
# We also support legacy setups where the model is in the text_encoders folder
if model_path is None:
model_path = folder_paths.get_full_path_or_raise("text_encoders", model_name)
model_path = folder_paths.get_full_path("text_encoders", model_name)
sd = load_torch_file(model_path, safe_load=True)
if "log_scale" not in sd:
raise ValueError("Invalid CLIP model, this node expectes the 'open-clip-xlm-roberta-large-vit-huge-14' model")
+39 -49
View File
@@ -296,6 +296,7 @@ class WanVideoSampler:
phantom_latents = fun_ref_image = ATI_tracks = None
add_cond = attn_cond = attn_cond_neg = noise_pred_flipped = None
humo_audio = humo_audio_neg = None
image_cond_mot_ref = None
#I2V
image_cond = image_embeds.get("image_embeds", None)
@@ -342,8 +343,7 @@ class WanVideoSampler:
dtype=torch.float32,
generator=seed_g,
device=torch.device("cpu"))
seq_len = math.ceil((noise.shape[2] * noise.shape[3]) / 4 * noise.shape[1])
seq_len = image_embeds["max_seq_len"]
control_embeds = image_embeds.get("control_embeds", None)
if control_embeds is not None:
@@ -412,7 +412,7 @@ class WanVideoSampler:
dtype=torch.float32,
device=torch.device("cpu"),
generator=seed_g)
seq_len = math.ceil((noise.shape[2] * noise.shape[3]) / 4 * noise.shape[1])
recammaster = image_embeds.get("recammaster", None)
@@ -484,13 +484,22 @@ class WanVideoSampler:
phantom_start_percent = image_embeds.get("phantom_start_percent", 0.0)
phantom_end_percent = image_embeds.get("phantom_end_percent", 1.0)
# Video-as-prompt (VAP)
mot_ref_clip_embeds = mot_ref_context = x_mot_ref = None
video_prompt_embeds = image_embeds.get("video_prompt_embeds", None)
if video_prompt_embeds is not None:
image_cond_mot_ref = video_prompt_embeds.get("image_embeds", None)
image_cond_mask_ = video_prompt_embeds.get("mask", None)
if image_cond_mask_ is not None:
image_cond_mot_ref = torch.cat([image_cond_mask_, image_cond_mot_ref])
latents_mot_ref = video_prompt_embeds.get("video_prompt_latents", None)
x_mot_ref = torch.cat([latents_mot_ref, image_cond_mot_ref], dim=0)
mot_ref_context = video_prompt_embeds.get("text_embeds", None)
mot_ref_clip_embeds = video_prompt_embeds.get("clip_context", None)
# CLIP image features
clip_fea = image_embeds.get("clip_context", None)
if clip_fea is not None:
clip_fea = clip_fea.to(dtype)
clip_fea_neg = image_embeds.get("negative_clip_context", None)
if clip_fea_neg is not None:
clip_fea_neg = clip_fea_neg.to(dtype)
num_frames = image_embeds.get("num_frames", 0)
@@ -916,9 +925,6 @@ class WanVideoSampler:
rope_function = "default" #echoshot does not support comfy rope function
log.info(f"Number of shots in prompt: {shot_num}, Shot token lengths: {shot_len}")
# Bindweave
qwenvl_embeds_pos = image_embeds.get("qwenvl_embeds_pos", None)
qwenvl_embeds_neg = image_embeds.get("qwenvl_embeds_neg", None)
mm.unload_all_models()
mm.soft_empty_cache()
@@ -1044,7 +1050,7 @@ class WanVideoSampler:
if standin_input is not None:
rope_function = "comfy" # only works with this currently
freqs = None
freqs = freqs_mot_ref = None
transformer.rope_embedder.k = None
transformer.rope_embedder.num_frames = None
d = transformer.dim // transformer.num_heads
@@ -1058,15 +1064,19 @@ class WanVideoSampler:
rope_params_mocha(1024, d - 4 * (d // 6), L_test=latent_video_length, k=riflex_freq_index, start=-1),
rope_params_mocha(1024, 2 * (d // 6), start=-1),
rope_params_mocha(1024, 2 * (d // 6), start=-1)
],
dim=1)
], dim=1)
elif "default" in rope_function or bidirectional_sampling: # original RoPE
freqs = torch.cat([
rope_params(1024, d - 4 * (d // 6), L_test=latent_video_length, k=riflex_freq_index),
rope_params(1024, 2 * (d // 6)),
rope_params(1024, 2 * (d // 6))
],
dim=1)
], dim=1).to(device)
if x_mot_ref is not None:
freqs_mot_ref = torch.cat([
rope_params(1024, d - 4 * (d // 6), L_test=latent_video_length, k=riflex_freq_index, mot_ref_latent=x_mot_ref),
rope_params(1024, 2 * (d // 6)),
rope_params(1024, 2 * (d // 6))
], dim=1).to(device)
elif "comfy" in rope_function: # comfy's rope
transformer.rope_embedder.k = riflex_freq_index
transformer.rope_embedder.num_frames = latent_video_length
@@ -1141,16 +1151,6 @@ class WanVideoSampler:
lynx_embeds["ref_buffer_uncond"] = lynx_ref_buffer_uncond if not math.isclose(cfg[0], 1.0) else None
mm.soft_empty_cache()
# UniLumos
foreground_latents = image_embeds.get("foreground_latents", None)
if foreground_latents is not None:
log.info(f"UniLumos foreground latent input shape: {foreground_latents.shape}")
foreground_latents = foreground_latents.to(device, dtype)
background_latents = image_embeds.get("background_latents", None)
if background_latents is not None:
log.info(f"UniLumos background latent input shape: {background_latents.shape}")
background_latents = background_latents.to(device, dtype)
#region model pred
def predict_with_cfg(z, cfg_scale, positive_embeds, negative_embeds, timestep, idx, image_cond=None, clip_fea=None,
control_latents=None, vace_data=None, unianim_data=None, audio_proj=None, control_camera_latents=None,
@@ -1361,18 +1361,9 @@ class WanVideoSampler:
z = z * c_in
timestep = c_noise
if image_cond is not None:
self.noise_front_pad_num = image_cond_input.shape[1] - z.shape[1]
if self.noise_front_pad_num > 0:
pad = torch.zeros((z.shape[0], self.noise_front_pad_num, z.shape[2], z.shape[3]), dtype=z.dtype, device=z.device)
z = torch.concat([pad, z], dim=1)
nonlocal seq_len
seq_len = math.ceil((z.shape[2] * z.shape[3]) / 4 * z.shape[1])
else:
self.noise_front_pad_num = 0
if background_latents is not None or foreground_latents is not None:
z = torch.cat([z, foreground_latents.to(z), background_latents.to(z)], dim=0)
x_mot_ref_input = None
if image_cond_mot_ref is not None:
x_mot_ref_input = [x_mot_ref.to(z)]
base_params = {
'x': [z], # latent
@@ -1429,7 +1420,11 @@ class WanVideoSampler:
"ovi_negative_text_embeds": ovi_negative_text_embeds, # Audio latent model negative text embeds for Ovi
"flashvsr_LQ_latent": flashvsr_LQ_latent, # FlashVSR LQ latent for upsampling
"flashvsr_strength": flashvsr_strength, # FlashVSR strength
"num_cond_latents": len(all_indices) if transformer.is_longcat else None,
"num_cond_latents": len(all_indices) if transformer.is_longcat else None, # number of cond latents LongCat to separate attention
"x_mot_ref": x_mot_ref_input, # motion reference latents for VAP
"mot_ref_context": mot_ref_context if image_cond_mot_ref is not None else None, # motion reference context for VAP
"mot_ref_clip_embeds": mot_ref_clip_embeds, # motion reference clip features for VAP
"freqs_mot_ref": freqs_mot_ref, # motion reference RoPE freqs for VAP
}
batch_size = 1
@@ -1445,7 +1440,6 @@ class WanVideoSampler:
#conditional (positive) pass
if pos_latent is not None: # for humo
base_params['x'] = [torch.cat([z[:, :-humo_reference_count], pos_latent], dim=1)]
base_params["add_text_emb"] = qwenvl_embeds_pos.to(device) if qwenvl_embeds_pos is not None else None # QwenVL embeddings for Bindweave
noise_pred_cond, noise_pred_ovi, cache_state_cond = transformer(
context=positive_embeds,
pred_id=cache_state[0] if cache_state else None,
@@ -1462,7 +1456,6 @@ class WanVideoSampler:
#unconditional (negative) pass
base_params['is_uncond'] = True
base_params['clip_fea'] = clip_fea_neg if clip_fea_neg is not None else clip_fea
base_params["add_text_emb"] = qwenvl_embeds_neg.to(device) if qwenvl_embeds_neg is not None else None # QwenVL embeddings for Bindweave
if wananim_face_pixels is not None:
base_params['wananim_face_pixel_values'] = torch.zeros_like(wananim_face_pixels).to(device, torch.float32) - 1
if humo_audio_input_neg is not None:
@@ -1586,23 +1579,23 @@ class WanVideoSampler:
noise_pred_uncond.view(batch_size, -1)
).view(batch_size, 1, 1, 1)
noise_pred_uncond_scaled = noise_pred_uncond * alpha
noise_pred_uncond = noise_pred_uncond * alpha
if use_tangential:
noise_pred_uncond_scaled = tangential_projection(noise_pred_cond, noise_pred_uncond_scaled)
noise_pred_uncond = tangential_projection(noise_pred_cond, noise_pred_uncond)
# RAAG (RATIO-aware Adaptive Guidance)
if raag_alpha > 0.0:
cfg_scale = get_raag_guidance(noise_pred_cond, noise_pred_uncond_scaled, cfg_scale, raag_alpha)
cfg_scale = get_raag_guidance(noise_pred_cond, noise_pred_uncond, cfg_scale, raag_alpha)
log.info(f"RAAG modified cfg: {cfg_scale}")
#https://github.com/WikiChao/FreSca
if use_fresca:
filtered_cond = fourier_filter(noise_pred_cond - noise_pred_uncond, fresca_scale_low, fresca_scale_high, fresca_freq_cutoff)
noise_pred = noise_pred_uncond_scaled + cfg_scale * filtered_cond * alpha
noise_pred = noise_pred_uncond + cfg_scale * filtered_cond * alpha
else:
noise_pred = noise_pred_uncond_scaled + cfg_scale * (noise_pred_cond - noise_pred_uncond_scaled)
del noise_pred_uncond_scaled, noise_pred_cond, noise_pred_uncond
noise_pred = noise_pred_uncond + cfg_scale * (noise_pred_cond - noise_pred_uncond)
del noise_pred_uncond, noise_pred_cond
if latent_model_input_ovi is not None:
if ovi_audio_cfg is None:
@@ -2991,9 +2984,6 @@ class WanVideoSampler:
if flowedit_args is None:
latent = latent.to(intermediate_device)
if self.noise_front_pad_num > 0:
noise_pred = noise_pred[:, self.noise_front_pad_num:]
if use_tsr:
noise_pred = temporal_score_rescaling(noise_pred, latent, timestep, tsr_k, tsr_sigma)
-97
View File
@@ -1,8 +1,6 @@
import torch
import numpy as np
from comfy.utils import common_upscale
from comfy import model_management
from tqdm import tqdm
from .utils import log
from einops import rearrange
@@ -14,9 +12,6 @@ except:
VAE_STRIDE = (4, 8, 8)
PATCH_SIZE = (1, 2, 2)
main_device = model_management.get_torch_device()
offload_device = model_management.unet_offload_device()
class WanVideoImageResizeToClosest:
@classmethod
def INPUT_TYPES(s):
@@ -665,96 +660,6 @@ class FaceMaskFromPoseKeypoints:
cv2.fillPoly(canvas, pts=[outer_contour], color=part_color)
return canvas
class DrawGaussianNoiseOnImage:
@classmethod
def INPUT_TYPES(s):
return {"required": {
"image": ("IMAGE", ),
"mask": ("MASK", ),
},
"optional": {
"device": (["cpu", "gpu"], {"default": "cpu", "tooltip": "Device to use for processing"}),
"seed": ("INT", {"default": 0, "min": 0, "max": 0xffffffffffffffff}),
}
}
RETURN_TYPES = ("IMAGE", )
RETURN_NAMES = ("images",)
FUNCTION = "apply"
CATEGORY = "KJNodes/masking"
DESCRIPTION = "Fills the background (masked area) with Gaussian noise sampled using the mean and variance of the subject (unmasked) region."
def apply(self, image, mask, device="cpu", seed=0):
B, H, W, C = image.shape
BM, HM, WM = mask.shape
processing_device = main_device if device == "gpu" else torch.device("cpu")
in_masks = mask.clone().to(processing_device)
in_images = image.clone().to(processing_device)
# Resize mask to match image dimensions
if HM != H or WM != W:
in_masks = F.interpolate(mask.unsqueeze(1), size=(H, W), mode='nearest-exact').squeeze(1)
# Match batch sizes
if B > BM:
in_masks = in_masks.repeat((B + BM - 1) // BM, 1, 1)[:B]
elif BM > B:
in_masks = in_masks[:B]
output_images = []
# Set random seed for reproducibility
generator = torch.Generator(device=processing_device).manual_seed(seed)
for i in tqdm(range(B), desc="DrawGaussianNoiseOnImage batch"):
curr_mask = in_masks[i]
img_idx = min(i, B - 1)
curr_image = in_images[img_idx]
# Expand mask to 3 channels
mask_expanded = curr_mask.unsqueeze(-1).expand(-1, -1, 3)
# Calculate mean and std per channel from the subject region (where mask is 1)
subject_mask = mask_expanded > 0.5
# Initialize noise tensor
noise = torch.zeros_like(curr_image)
for c in range(C):
channel = curr_image[:, :, c]
channel_mask = subject_mask[:, :, c]
if channel_mask.sum() > 0:
# Get subject pixels
subject_pixels = channel[channel_mask]
# Calculate statistics
mean = subject_pixels.mean()
std = subject_pixels.std()
# Generate Gaussian noise for this channel
noise[:, :, c] = torch.normal(mean=mean.item(), std=std.item(),
size=(H, W), generator=generator,
device=processing_device)
# Clamp noise to valid range
noise = torch.clamp(noise, 0.0, 1.0)
# Apply: keep subject, fill background with noise
masked_image = curr_image * mask_expanded + noise * (1 - mask_expanded)
output_images.append(masked_image)
# If no masks were processed, return empty tensor
if not output_images:
return (torch.zeros((0, H, W, 3), dtype=image.dtype),)
out_rgb = torch.stack(output_images, dim=0).cpu()
return (out_rgb, )
NODE_CLASS_MAPPINGS = {
"WanVideoImageResizeToClosest": WanVideoImageResizeToClosest,
@@ -768,7 +673,6 @@ NODE_CLASS_MAPPINGS = {
"NormalizeAudioLoudness": NormalizeAudioLoudness,
"WanVideoPassImagesFromSamples": WanVideoPassImagesFromSamples,
"FaceMaskFromPoseKeypoints": FaceMaskFromPoseKeypoints,
"DrawGaussianNoiseOnImage": DrawGaussianNoiseOnImage,
}
NODE_DISPLAY_NAME_MAPPINGS = {
"WanVideoImageResizeToClosest": "WanVideo Image Resize To Closest",
@@ -782,5 +686,4 @@ NODE_DISPLAY_NAME_MAPPINGS = {
"NormalizeAudioLoudness": "Normalize Audio Loudness",
"WanVideoPassImagesFromSamples": "WanVideo Pass Images From Samples",
"FaceMaskFromPoseKeypoints": "Face Mask From Pose Keypoints",
"DrawGaussianNoiseOnImage": "Draw Gaussian Noise On Image",
}
+1 -1
View File
@@ -1,7 +1,7 @@
[project]
name = "ComfyUI-WanVideoWrapper"
description = "ComfyUI wrapper nodes for WanVideo"
version = "1.3.9"
version = "1.3.8"
license = {file = "LICENSE"}
dependencies = ["accelerate >= 1.2.1", "diffusers >= 0.33.0", "peft >= 0.17.0", "ftfy", "gguf >= 0.17.1", "pyloudnorm"]
+2 -24
View File
@@ -1,29 +1,7 @@
# ComfyUI wrapper nodes for [WanVideo](https://github.com/Wan-Video/Wan2.1) and related models.
## Update notification that can affect memory use in old workflows
In a recent update I changed how unmerged LoRA weights are handled:
Previously mostly due to my laziness they were always loaded from RAM when used, this was of course inefficient and also made using torch.compile for LoRA applying difficult, thus forcing a graph break when using unmerged LoRAs.
Now the LoRA weights are assigned as buffers to the corresponding modules, so they are part of the blocks and obey the block swapping unifying the offloading and allowing LoRA weights to benefit from the prefetch feature for async offoading. Downside is that this means if you did not use block swap, you will see increased memory use as the LoRAs are part of the model and all on VRAM.
If you use block swap, the LoRAs are swapped along the rest of the block, but the block size is now larger, this means you may have to compensate with couple of more blocks swapped.
Example situation: you use 1GB LoRA unmerged and swap 20 blocks on 14B model, we can divide the LoRA size by block count, single block grows by 25MB, 20 blocks grow by 500MB, so your VRAM usage would be 500MB more than before, to compensate you swap 2 more blocks.
### Unrelated other VRAM issue with torch.compile
After any update that modifies the model code and when using torch.compile it's common to run into issues with VRAM, this can be caused by using older pytorch/triton version without latest compile fixes, and/or from old triton caches, mostly in Windows. This manifests in the issue that first run of new input size may have drastically increased memory use, which can clear from simply running it again, and once cached, not manifest again. Again I've only seen this happen in Windows.
To clear your Triton cache you can delete the contents of following (default) folders:
`C:\Users\<username>\.triton`
`C:\Users\<username>\AppData\Local\Temp\torchinductor_<username>`
## Note: Due to the stupid amount of bots or people thinking this is some of video generation service, I've blocked new accounts from posting issues for now.
# ComfyUI wrapper nodes for [WanVideo](https://github.com/Wan-Video/Wan2.1) and related models.
# WORK IN PROGRESS (perpetually)
# Why should I use custom nodes when WanVideo works natively?
+183 -76
View File
@@ -249,7 +249,7 @@ def sinusoidal_embedding_1d(dim, position):
x = torch.cat([torch.cos(sinusoid), torch.sin(sinusoid)], dim=1)
return x
def rope_params(max_seq_len, dim, theta=10000, L_test=25, k=0, freqs_scaling=1.0):
def rope_params(max_seq_len, dim, theta=10000, L_test=25, k=0, freqs_scaling=1.0, mot_ref_latent=None):
assert dim % 2 == 0
exponents = torch.arange(0, dim, 2, dtype=torch.float64).div(dim)
inv_theta_pow = 1.0 / torch.pow(theta, exponents)
@@ -259,9 +259,16 @@ def rope_params(max_seq_len, dim, theta=10000, L_test=25, k=0, freqs_scaling=1.0
inv_theta_pow[k-1] = 0.9 * 2 * torch.pi / L_test
inv_theta_pow *= freqs_scaling
if mot_ref_latent is not None:
freqs = torch.arange(-mot_ref_latent.shape[1], max_seq_len)
freqs = torch.outer(freqs, inv_theta_pow)
freqs = torch.polar(torch.ones_like(freqs), freqs)
freqs = freqs[:max_seq_len]
freqs = torch.outer(torch.arange(max_seq_len), inv_theta_pow)
freqs = torch.polar(torch.ones_like(freqs), freqs)
else:
freqs = torch.outer(torch.arange(max_seq_len), inv_theta_pow)
freqs = torch.polar(torch.ones_like(freqs), freqs)
return freqs
@torch.autocast(device_type=mm.get_autocast_device(mm.get_torch_device()), enabled=False)
@@ -356,10 +363,10 @@ class WanRMSNorm(nn.Module):
if use_chunked:
return self.forward_chunked(x, num_chunks)
else:
return self._norm(x.to(self.weight.dtype)) * self.weight
return (self._norm(x.to(self.weight.dtype)) * self.weight).to(x.dtype)
def _norm(self, x):
return x * (torch.rsqrt(x.pow(2).mean(dim=-1, keepdim=True) + self.eps)).to(x.dtype)
return x * (torch.rsqrt(x.pow(2).mean(dim=-1, keepdim=True) + self.eps))
def forward_chunked(self, x, num_chunks=4):
output = torch.empty_like(x)
@@ -386,7 +393,7 @@ class WanFusedRMSNorm(nn.RMSNorm):
if use_chunked:
return self.forward_chunked(x, num_chunks)
else:
return super().forward(x)
return super().forward(x.to(self.weight.dtype).to(x.dtype))
def forward_chunked(self, x, num_chunks=4):
output = torch.empty_like(x)
@@ -398,7 +405,7 @@ class WanFusedRMSNorm(nn.RMSNorm):
for size in chunk_sizes:
end_idx = start_idx + size
chunk = x[:, start_idx:end_idx, :]
output[:, start_idx:end_idx, :] = super().forward(chunk)
output[:, start_idx:end_idx, :] = super().forward(chunk.to(self.weight.dtype)).to(chunk.dtype)
start_idx = end_idx
return output
@@ -464,8 +471,8 @@ class WanSelfAttention(nn.Module):
def qkv_fn(self, x):
b, s, n, d = *x.shape[:2], self.num_heads, self.head_dim
q = self.norm_q(self.q(x).to(self.norm_q.weight.dtype)).to(x.dtype).view(b, s, n, d)
k = self.norm_k(self.k(x).to(self.norm_k.weight.dtype)).to(x.dtype).view(b, s, n, d)
q = self.norm_q(self.q(x)).view(b, s, n, d)
k = self.norm_k(self.k(x)).view(b, s, n, d)
v = self.v(x).view(b, s, n, d)
return q, k, v
@@ -480,8 +487,8 @@ class WanSelfAttention(nn.Module):
def qkv_fn_ip(self, x):
b, s, n, d = *x.shape[:2], self.num_heads, self.head_dim
q = self.norm_q(self.q(x) + self.q_loras(x).to(self.norm_q.weight.dtype)).to(x.dtype).view(b, s, n, d)
k = self.norm_k(self.k(x) + self.k_loras(x).to(self.norm_k.weight.dtype)).to(x.dtype).view(b, s, n, d)
q = self.norm_q(self.q(x) + self.q_loras(x)).view(b, s, n, d)
k = self.norm_k(self.k(x) + self.k_loras(x)).view(b, s, n, d)
v = (self.v(x) + self.v_loras(x)).view(b, s, n, d)
return q, k, v
@@ -674,7 +681,7 @@ class WanT2VCrossAttention(WanSelfAttention):
if num_cond_latents is not None and num_cond_latents > 0:
num_cond_latents_thw = num_cond_latents * (s // num_latent_frames)
x = x[:, num_cond_latents_thw:]
q = self.norm_q(self.q(x).view(b, -1, n, d))
q = self.norm_q(self.q(x).view(b, -1, n, d).to(self.norm_q.weight.dtype)).to(x.dtype)
else:
q = self.norm_q(self.q(x).to(self.norm_q.weight.dtype),num_chunks=2 if rope_func == "comfy_chunked" else 1).to(x.dtype).view(b, -1, n, d)
@@ -747,6 +754,32 @@ class WanT2VCrossAttention(WanSelfAttention):
return torch.cat([torch.zeros((b, num_cond_latents_thw, x.shape[-1]), dtype=x.dtype, device=x.device), self.o(x)], dim=1).contiguous()
return self.o(x)
class WanCrossAttentionMOTRef(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", head_norm=False):
super().__init__(in_features, out_features, num_heads, qk_norm, eps, kv_dim=kv_dim, rms_norm_function=rms_norm_function, head_norm=head_norm)
self.k_img = nn.Linear(in_features, out_features)
self.v_img = nn.Linear(in_features, out_features)
self.norm_k_img = WanRMSNorm(out_features, eps=eps) if qk_norm else nn.Identity()
self.attention_mode = attention_mode
def forward(self, x, context, grid_sizes=None, clip_embed=None, rope_func="comfy", **kwargs):
b, n, d = x.size(0), self.num_heads, self.head_dim
q = self.norm_q(self.q(x).to(self.norm_q.weight.dtype), num_chunks=2 if rope_func == "comfy_chunked" else 1).to(x.dtype).view(b, -1, n, d)
k = self.norm_k(self.k(context).to(self.norm_k.weight.dtype)).to(x.dtype).view(b, -1, n, d)
v = self.v(context).view(b, -1, n, d)
x = attention(q, k, v, attention_mode=self.attention_mode).flatten(2)
if clip_embed is not None:
k_img = self.norm_k_img(self.k_img(clip_embed)).view(b, -1, n, d)
v_img = self.v_img(clip_embed).view(b, -1, n, d)
img_x = attention(q, k_img, v_img, attention_mode=self.attention_mode).flatten(2)
x = x + img_x
return self.o(x)
class WanI2VCrossAttention(WanSelfAttention):
@@ -773,13 +806,13 @@ class WanI2VCrossAttention(WanSelfAttention):
x_text = self.normalized_attention_guidance(b, n, d, q, context, nag_context, nag_params)
else:
# text attention
k = self.norm_k(self.k(context).to(self.norm_k.weight.dtype)).view(b, -1, n, d).to(x.dtype)
k = self.norm_k(self.k(context)).view(b, -1, n, d)
v = self.v(context).view(b, -1, n, d)
x_text = attention(q, k, v, attention_mode=self.attention_mode).flatten(2)
#img attention
if clip_embed is not None:
k_img = self.norm_k_img(self.k_img(clip_embed).to(self.norm_k_img.weight.dtype)).view(b, -1, n, d).to(x.dtype)
k_img = self.norm_k_img(self.k_img(clip_embed)).view(b, -1, n, d)
v_img = self.v_img(clip_embed).view(b, -1, n, d)
img_x = attention(q, k_img, v_img, attention_mode=self.attention_mode).flatten(2)
x = x_text + img_x
@@ -871,8 +904,8 @@ class MTVCrafterMotionAttention(WanSelfAttention):
b, n, d = x.size(0), self.num_heads, self.head_dim
# compute query, key, value
q = self.norm_q(self.q(x)).view(b, -1, n, d)
k = self.norm_k(self.k(mo)).view(b, n, -1, d)
q = self.norm_q(self.q(x).to(self.norm_q.weight.dtype)).to(x.dtype).view(b, -1, n, d)
k = self.norm_k(self.k(mo).to(self.norm_k.weight.dtype)).to(x.dtype).view(b, n, -1, d)
v = self.v(mo).view(b, -1, n, d)
# compute attention
@@ -897,7 +930,7 @@ class WanAttentionBlock(nn.Module):
cross_attn_type, in_features, out_features, ffn_dim, ffn2_dim, num_heads,
qk_norm=True, cross_attn_norm=False, eps=1e-6, attention_mode="sdpa", rope_func="comfy", rms_norm_function="default",
use_motion_attn=False, use_humo_audio_attn=False, face_fuser_block=False, lynx_ip_layers=None, lynx_ref_layers=None,
block_idx=0, is_longcat=False):
block_idx=0, mot_ref_block=False, is_longcat=False):
super().__init__()
self.dim = out_features
self.ffn_dim = ffn_dim
@@ -913,6 +946,7 @@ class WanAttentionBlock(nn.Module):
self.dense_block = False
self.dense_attention_mode = "sageattn"
self.block_idx = block_idx
self.mot_ref_block = mot_ref_block
self.kv_cache = None
self.use_motion_attn = use_motion_attn
@@ -952,6 +986,16 @@ class WanAttentionBlock(nn.Module):
self.seg_idx = None
# video-as-prompt (VAP)
if mot_ref_block:
self.norm1_mot_ref = WanLayerNorm(self.dim, eps)
self.norm2_mot_ref = WanLayerNorm(self.dim, eps)
self.norm3_mot_ref = WanLayerNorm(out_features, eps, elementwise_affine=True) if cross_attn_norm else nn.Identity()
self.self_attn_mot_ref = WanSelfAttention(in_features, out_features, num_heads, qk_norm, eps, self.attention_mode, rms_norm_function=rms_norm_function)
self.cross_attn_mot_ref = WanCrossAttentionMOTRef(in_features, out_features, num_heads, qk_norm, eps, rms_norm_function=rms_norm_function)
self.modulation_mot_ref = nn.Parameter(torch.randn(1, 6, out_features) / in_features**0.5)
self.ffn_mot_ref = nn.Sequential(nn.Linear(in_features, ffn_dim), nn.GELU(approximate='tanh'), nn.Linear(ffn2_dim, out_features))
# HuMo audio cross-attn
if use_humo_audio_attn:
self.audio_cross_attn_wrapper = AudioCrossAttentionWrapper(in_features, out_features, num_heads, qk_norm, eps, kv_dim=1536)
@@ -1045,7 +1089,8 @@ class WanAttentionBlock(nn.Module):
lynx_x_ip=None, lynx_ref_feature=None, lynx_ip_scale=1.0, lynx_ref_scale=1.0, #lynx
x_ovi=None, e_ovi=None, freqs_ovi=None, context_ovi=None, seq_lens_ovi=None, grid_sizes_ovi=None,
num_cond_latents=None, #longcat image cond amount
):
# VAP
x_mot_ref=None, context_mot_ref=None, grid_sizes_mot_ref=None, e_mot_ref=None, freqs_mot_ref=None, clip_embed_mot_ref=None, num_mot_ref=1):
r"""
Args:
x(Tensor): Shape [B, L, C]
@@ -1081,6 +1126,12 @@ class WanAttentionBlock(nn.Module):
input_x = torch.concat([input_x, input_x_ip], dim=1)
self.kv_cache = None
# video-as-prompt motion reference
use_mot_ref = x_mot_ref is not None and self.mot_ref_block
if use_mot_ref:
shift_msa_mot_ref, scale_msa_mot_ref, gate_msa_mot_ref, shift_mlp_mot_ref, scale_mlp_mot_ref, gate_mlp_mot_ref = self.get_mod(e_mot_ref.to(x.device), self.modulation_mot_ref)
input_x_mot_ref = self.modulate(self.norm1_mot_ref(x_mot_ref.to(shift_msa_mot_ref.dtype)), shift_msa_mot_ref, scale_msa_mot_ref).to(input_dtype)
if x_ovi is not None:
shift_msa_ovi, scale_msa_ovi, gate_msa_ovi, shift_mlp_ovi, scale_mlp_ovi, gate_mlp_ovi = self.get_mod(e_ovi.to(x.device), self.audio_block.modulation)
input_x_ovi = self.modulate(self.audio_block.norm1(x_ovi), shift_msa_ovi, scale_msa_ovi)
@@ -1125,8 +1176,13 @@ class WanAttentionBlock(nn.Module):
q, k, v = self.self_attn.qkv_fn_longcat(input_x)
else:
q, k, v = self.self_attn.qkv_fn(input_x)
if use_mot_ref:
q_mot_ref, k_mot_ref, v_mot_ref = self.self_attn_mot_ref.qkv_fn(input_x_mot_ref)
# Apply RoPE
if self.rope_func == "comfy":
q, k = apply_rope_comfy(q, k, freqs)
if use_mot_ref:
q_mot_ref, k_mot_ref = apply_rope_comfy(q_mot_ref, k_mot_ref, freqs_mot_ref)
elif self.rope_func == "comfy_chunked":
q, k = apply_rope_comfy_chunked(q, k, freqs)
elif self.rope_func == "mocha":
@@ -1136,6 +1192,9 @@ class WanAttentionBlock(nn.Module):
else:
q = rope_apply(q, grid_sizes, freqs, reverse_time=reverse_time)
k = rope_apply(k, grid_sizes, freqs, reverse_time=reverse_time)
if use_mot_ref:
q_mot_ref = rope_apply(q_mot_ref, grid_sizes_mot_ref, freqs_mot_ref)
k_mot_ref = rope_apply(k_mot_ref, grid_sizes_mot_ref, freqs_mot_ref)
if x_ovi is not None:
q_ovi, k_ovi, v_ovi = self.audio_block.self_attn.qkv_fn(input_x_ovi)
@@ -1196,6 +1255,20 @@ class WanAttentionBlock(nn.Module):
x_noise = self.self_attn.forward(q[:, num_cond_latents_thw:].contiguous(), k, v, seq_lens)
# merge x_cond and x_noise
y = torch.cat([x_cond, x_noise], dim=1).contiguous()
elif use_mot_ref:
y_temp = attention(
torch.cat([q, q_mot_ref], dim=1),
torch.cat([k, k_mot_ref], dim=1),
torch.cat([v, v_mot_ref], dim=1),
attention_mode=self.attention_mode
)
y, y_mot_ref = (
y_temp[:, :q.shape[1]],
y_temp[:, q.shape[1]:q.shape[1]+q_mot_ref.shape[1]]
)
y = self.self_attn.o(y.flatten(2))
y_mot_ref = self.self_attn_mot_ref.o(y_mot_ref.flatten(2))
del y_temp
else:
y = self.self_attn.forward(q, k, v, seq_lens, lynx_ref_feature=lynx_ref_feature, lynx_ref_scale=lynx_ref_scale)
@@ -1227,8 +1300,11 @@ class WanAttentionBlock(nn.Module):
else:
if not is_longcat:
x = x.addcmul(y, gate_msa)
if use_mot_ref:
x_mot_ref = x_mot_ref.addcmul(y_mot_ref, gate_msa_mot_ref)
else:
x = x + (y.view(B, -1, N//T, C).float() * gate_msa).to(input_dtype).view(B, -1, C)
del y, gate_msa
# cross-attention & ffn function
@@ -1263,6 +1339,10 @@ class WanAttentionBlock(nn.Module):
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=original_seq_len, lynx_x_ip=lynx_x_ip, lynx_ip_scale=lynx_ip_scale, num_cond_latents=num_cond_latents)
x = x.to(input_dtype)
if use_mot_ref:
x_mot_ref = x_mot_ref + self.cross_attn_mot_ref(self.norm3_mot_ref(x_mot_ref.to(self.norm3_mot_ref.weight.dtype)).to(input_dtype), context_mot_ref, grid_sizes_mot_ref,
clip_embed=clip_embed_mot_ref)
x_mot_ref = x_mot_ref.to(input_dtype)
# MultiTalk
if multitalk_audio_embedding is not None and not isinstance(self, VaceWanAttentionBlock):
x_audio = self.audio_cross_attn(self.norm_x(x.to(self.norm_x.weight.dtype)).to(input_dtype), encoder_hidden_states=multitalk_audio_embedding,
@@ -1270,8 +1350,8 @@ class WanAttentionBlock(nn.Module):
x = x.add(x_audio, alpha=audio_scale)
# MTV-Crafter Motion Attention
if self.use_motion_attn and mtv_motion_tokens is not None and mtv_motion_rotary_emb is not None:
x_motion = self.motion_attn(self.norm4(x), mtv_motion_tokens, mtv_motion_rotary_emb, grid_sizes, mtv_freqs)
if self.use_motion_attn and mtv_motion_tokens is not None and mtv_motion_rotary_emb is not None:
x_motion = self.motion_attn(self.norm4(x.to(self.norm4.weight.dtype)).to(input_dtype), mtv_motion_tokens, mtv_motion_rotary_emb, grid_sizes, mtv_freqs)
x = x.add(x_motion, alpha=mtv_strength)
# HuMo Audio Cross-Attention
@@ -1317,7 +1397,13 @@ class WanAttentionBlock(nn.Module):
x_ip = x_ip.addcmul(y_ip, gate_msa_ip)
y_ip = self.ffn(torch.addcmul(shift_mlp_ip, self.norm2(x_ip), 1 + scale_mlp_ip))
x_ip = x_ip.addcmul(y_ip, gate_mlp_ip)
return x, x_ip, lynx_ref_feature, x_ovi
if use_mot_ref:
norm2_x_mot_ref = self.norm2_mot_ref(x_mot_ref.to(shift_mlp_mot_ref.dtype))
mod_x_mot_ref = torch.addcmul(shift_mlp_mot_ref, norm2_x_mot_ref, 1 + scale_mlp_mot_ref)
x_ffn_mot_ref = self.ffn_mot_ref(mod_x_mot_ref.to(input_dtype))
x_mot_ref = x_mot_ref.addcmul(x_ffn_mot_ref, gate_mlp_mot_ref)
return x, x_ip, lynx_ref_feature, x_ovi, x_mot_ref
@torch.compiler.disable()
def split_cross_attn_ffn(self, x, context, shift_mlp, scale_mlp, gate_mlp, clip_embed=None, grid_sizes=None):
@@ -1517,10 +1603,10 @@ class MLPProj(torch.nn.Module):
if fl_pos_emb: # NOTE: we only use this for `fl2v`
self.emb_pos = nn.Parameter(torch.zeros(1, 257 * 2, 1280))
def forward(self, image_embeds):
def forward(self, image_embeds, dtype=torch.float32):
if hasattr(self, 'emb_pos'):
image_embeds = image_embeds + self.emb_pos.to(image_embeds.device)
clip_extra_context_tokens = self.proj(image_embeds)
clip_extra_context_tokens = self.proj(image_embeds.to(self.proj[1].weight.dtype)).to(dtype)
return clip_extra_context_tokens
from .s2v.auxi_blocks import MotionEncoder_tc
@@ -1622,47 +1708,22 @@ class AudioInjector_WAN(nn.Module):
class WanModel(torch.nn.Module):
def __init__(self,
model_type='t2v',
patch_size=(1, 2, 2),
text_len=512,
in_dim=16,
dim=2048,
in_features=5120,
out_features=5120,
ffn_dim=8192,
ffn2_dim=8192,
freq_dim=256,
text_dim=4096,
out_dim=16,
num_heads=16,
num_layers=32,
qk_norm=True,
cross_attn_norm=True,
eps=1e-6,
attention_mode='sdpa',
rope_func='comfy',
rms_norm_function='default',
main_device=torch.device('cuda'),
offload_device=torch.device('cpu'),
model_type='t2v', patch_size=(1, 2, 2),
text_len=512, in_dim=16, dim=2048, in_features=5120, out_features=5120,
ffn_dim=8192, ffn2_dim=8192, freq_dim=256, text_dim=4096, out_dim=16,
num_heads=16, num_layers=32, qk_norm=True, cross_attn_norm=True,
eps=1e-6, attention_mode='sdpa', rope_func='comfy', rms_norm_function='default',
main_device=torch.device('cuda'), offload_device=torch.device('cpu'),
dtype=torch.float16,
teacache_coefficients=[],
magcache_ratios=[],
vace_layers=None,
vace_in_dim=None,
inject_sample_info=False,
add_ref_conv=False,
in_dim_ref_conv=16,
add_control_adapter=False,
in_dim_control_adapter=24,
use_motion_attn=False,
teacache_coefficients=[], magcache_ratios=[],
vace_layers=None, vace_in_dim=None,
inject_sample_info=False, add_ref_conv=False,
in_dim_ref_conv=16, add_control_adapter=False, in_dim_control_adapter=24, use_motion_attn=False,
#s2v
cond_dim=0,
audio_dim=1024,
num_audio_token=4,
enable_adain=False,
adain_mode="attn_norm",
cond_dim=0, audio_dim=1024, num_audio_token=4, enable_adain=False, adain_mode="attn_norm",
audio_inject_layers=[0, 4, 8, 12, 16, 20, 24, 27, 30, 33, 36, 39],
zero_timestep=False,
# humo
humo_audio=False,
# WanAnimate
is_wananimate=False,
@@ -1670,8 +1731,8 @@ class WanModel(torch.nn.Module):
# lynx
lynx_ip_layers=None,
lynx_ref_layers=None,
# ovi
is_ovi_audio_model=False,
# VAP
is_VAP = False,
# LongCat
is_longcat=False,
):
@@ -1808,7 +1869,7 @@ class WanModel(torch.nn.Module):
nn.SiLU(),
ConvMLP(dim, dim * 4, kernel_size=7, padding=3),
)
self.original_patch_embedding = self.patch_embedding
self.expanded_patch_embedding = self.patch_embedding
@@ -1825,6 +1886,13 @@ class WanModel(torch.nn.Module):
adaln_tembed_dim = 512
self.time_embedding = TimestepEmbedder(t_embed_dim=adaln_tembed_dim, frequency_embedding_size=freq_dim)
if is_VAP:
self.patch_embedding_mot_ref = nn.Conv3d(in_dim, dim, kernel_size=patch_size, stride=patch_size)
self.time_embedding_mot_ref = nn.Sequential(nn.Linear(freq_dim, dim), nn.SiLU(), nn.Linear(dim, dim))
self.time_projection_mot_ref = nn.Sequential(nn.SiLU(), nn.Linear(dim, dim * 6))
self.text_embedding_mot_ref = nn.Sequential(nn.Linear(text_dim, dim), nn.GELU(approximate='tanh'), nn.Linear(dim, dim))
self.img_emb_mot_ref = MLPProj(1280, dim)
VAP_layers = [0, 4, 8, 12, 16, 20, 24, 28, 32, 36]
if vace_layers is not None:
self.vace_layers = [i for i in range(0, self.num_layers, 2)] if vace_layers is None else vace_layers
@@ -1865,7 +1933,7 @@ class WanModel(torch.nn.Module):
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), lynx_ip_layers=lynx_ip_layers, lynx_ref_layers=lynx_ref_layers,
block_idx=i, is_longcat=is_longcat)
block_idx=i, is_longcat=is_longcat, mot_ref_block=is_VAP and i in VAP_layers)
for i in range(num_layers)
])
#MTV Crafter
@@ -2139,12 +2207,15 @@ class WanModel(torch.nn.Module):
return x.add(residual_out, alpha=strength)
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):
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, mot_ref=False):
patch_size = self.patch_size
t_len = ((t + (patch_size[0] // 2)) // patch_size[0])
h_len = ((h + (patch_size[1] // 2)) // patch_size[1])
w_len = ((w + (patch_size[2] // 2)) // patch_size[2])
if mot_ref:
t_start = -t_len
if steps_t is None:
steps_t = t_len
if steps_h is None:
@@ -2226,7 +2297,7 @@ class WanModel(torch.nn.Module):
x_ovi=None, seq_len_ovi=None, ovi_negative_text_embeds=None,
flashvsr_LQ_latent=None, flashvsr_strength=1.0,
num_cond_latents=None,
add_text_emb=None,
x_mot_ref=None, mot_ref_context=None, mot_ref_clip_embeds=None, freqs_mot_ref=None,
):
r"""
Forward pass through the diffusion model
@@ -2257,6 +2328,7 @@ class WanModel(torch.nn.Module):
if mtv_motion_tokens is not None:
bs, motion_seq_len = mtv_motion_tokens.shape[0], mtv_motion_tokens.shape[1]
mtv_motion_tokens = torch.cat([mtv_motion_tokens, self.pad_motion_tokens.to(mtv_motion_tokens).expand(bs, motion_seq_len, -1)], dim=-1)
mtv_motion_tokens = mtv_motion_tokens.to(self.base_dtype)
# Fantasy Portrait
adapter_proj = ip_scale = None
@@ -2347,6 +2419,18 @@ class WanModel(torch.nn.Module):
d = self.dim // self.num_heads
freqs_ovi = rope_params(1024, d - 4 * (d // 6), freqs_scaling=0.19676).to(self.main_device)
x_ovi = x_ovi.to(self.main_device, self.base_dtype)
# video-as-prompt motion ref
if x_mot_ref is not None:
x_mot_ref = [self.patch_embedding_mot_ref(u.unsqueeze(0).to(torch.float32)).to(x[0].dtype) for u in x_mot_ref]
grid_sizes_mot_ref = torch.stack([torch.tensor(u.shape[2:], device=device, dtype=torch.long) for u in x_mot_ref])
x_mot_ref = [u.flatten(2).transpose(1, 2) for u in x_mot_ref]
seq_lens_mot_ref = torch.tensor([u.size(1) for u in x_mot_ref], dtype=torch.int32)
x_mot_ref = torch.cat([torch.cat([u, u.new_zeros(1, seq_lens_mot_ref - u.size(1), u.size(2))], dim=1) for u in x_mot_ref])
x_mot_ref = x_mot_ref.to(self.main_device, self.base_dtype)
num_mot_ref = 1
# WanAnimate
motion_vec = None
@@ -2457,13 +2541,16 @@ class WanModel(torch.nn.Module):
s2v_ref_latent.shape[4],
t_start=max(30, F + 9), device=x.device, dtype=x.dtype)
freqs = torch.cat([freqs, freqs_ref], dim=1)
self.cached_freqs = freqs
self.cached_shape = current_shape
self.cached_cond = has_cond
self.cached_rope_k = self.rope_embedder.k
self.cached_ntk_alphas = ntk_alphas
if x_mot_ref is not None:
freqs_mot_ref = self.rope_encode_comfy(F, H, W, mot_ref=True, freq_offset=freq_offset, ntk_alphas=ntk_alphas, attn_cond_shape=attn_cond_shape, device=x.device, dtype=x.dtype)
# Stand-In RoPE frequencies
if x_ip is not None:
# Generate RoPE frequencies for x_ip
@@ -2504,6 +2591,10 @@ class WanModel(torch.nn.Module):
time_embed_dtype = self.base_dtype
e = self.time_embedding(sinusoidal_embedding_1d(self.freq_dim, t.flatten()).to(time_embed_dtype)) # b, dim
e0 = self.time_projection(e).unflatten(1, (6, self.dim)) # b, 6, dim
if x_mot_ref is not None:
t_mod_ref = torch.tensor([1], device=t.device, dtype=t.dtype)
e_mot_ref = self.time_embedding_mot_ref(sinusoidal_embedding_1d(self.freq_dim, t_mod_ref.flatten()).to(time_embed_dtype)) # b, dim
e0_mot_ref = self.time_projection_mot_ref(e_mot_ref).unflatten(1, (6, self.dim)) # b, 6, dim
else:
time_embed_dtype = self.time_embedding.mlp[0].weight.dtype
if time_embed_dtype not in [torch.float16, torch.bfloat16, torch.float32]:
@@ -2568,13 +2659,20 @@ class WanModel(torch.nn.Module):
e = e.to(self.offload_device, non_blocking=self.use_non_blocking)
# clip vision embedding
clip_embed = None
clip_embed = clip_embed_mot_ref = None
if clip_fea is not None and hasattr(self, "img_emb"):
clip_fea = clip_fea.to(self.main_device)
if self.offload_img_emb:
self.img_emb.to(self.main_device)
clip_embed = self.img_emb(clip_fea) # bs x 257 x dim
#context = torch.concat([context_clip, context], dim=1)
if self.offload_img_emb:
self.img_emb.to(self.offload_device, non_blocking=self.use_non_blocking)
if mot_ref_clip_embeds is not None:
mot_ref_clip_embeds = mot_ref_clip_embeds.to(self.main_device)
if self.offload_img_emb:
self.img_emb.to(self.main_device)
clip_embed_mot_ref = self.img_emb_mot_ref(mot_ref_clip_embeds) # bs x 257 x dim
if self.offload_img_emb:
self.img_emb.to(self.offload_device, non_blocking=self.use_non_blocking)
@@ -2600,13 +2698,13 @@ class WanModel(torch.nn.Module):
torch.stack([torch.cat([u, u.new_zeros(self.text_len - u.size(0), u.size(1))]) for u in context_ovi]).to(text_embed_dtype))
tokens = context[0].shape[0]
context = torch.stack([torch.cat([u, u.new_zeros(self.text_len - u.size(0), u.size(1))]) for u in context]).to(text_embed_dtype)
context = self.text_embedding(
torch.stack([torch.cat([u, u.new_zeros(self.text_len - u.size(0), u.size(1))]) for u in context]).to(text_embed_dtype))
if add_text_emb is not None:
self.text_projection.to(self.main_device)
add_text_emb = self.text_projection(add_text_emb.to(self.text_projection[0].weight.dtype)).to(text_embed_dtype)
context = torch.cat([add_text_emb, context], dim=1)
context = self.text_embedding(context)
if mot_ref_context is not None:
context_mot_ref = mot_ref_context["prompt_embeds"] if not is_uncond else mot_ref_context["negative_prompt_embeds"]
context_mot_ref = self.text_embedding_mot_ref(
torch.stack([torch.cat([u, u.new_zeros(self.text_len - u.size(0), u.size(1))]) for u in context_mot_ref]).to(text_embed_dtype))
if self.is_longcat:
context[:, tokens:] = 0
@@ -2850,6 +2948,13 @@ class WanModel(torch.nn.Module):
kwargs['grid_sizes_ovi'] = grid_sizes_ovi
kwargs['seq_lens_ovi'] = seq_lens_ovi
kwargs['freqs_ovi'] = freqs_ovi
if x_mot_ref is not None:
kwargs['context_mot_ref'] = context_mot_ref
kwargs['freqs_mot_ref'] = freqs_mot_ref
kwargs['grid_sizes_mot_ref'] = grid_sizes_mot_ref
kwargs['e_mot_ref'] = e0_mot_ref.to(self.base_dtype)
kwargs['num_mot_ref'] = num_mot_ref
kwargs['clip_embed_mot_ref'] = clip_embed_mot_ref
if vace_data is not None:
@@ -2937,7 +3042,9 @@ class WanModel(torch.nn.Module):
if b in self.slg_blocks and is_uncond:
if self.slg_start_percent <= current_step_percentage <= self.slg_end_percent:
continue
x, x_ip, lynx_ref_feature, x_ovi = block(x, x_ip=x_ip, lynx_ref_feature=lynx_ref_feature, x_ovi=x_ovi, **kwargs) #run block
# ====run block start=====
x, x_ip, lynx_ref_feature, x_ovi, x_mot_ref = block(x, x_ip=x_ip, lynx_ref_feature=lynx_ref_feature, x_ovi=x_ovi, x_mot_ref=x_mot_ref, **kwargs)
# ====end 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 block.has_face_fuser_block and motion_vec is not None: