Compare commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
f685ee33ac | ||
|
|
4e31081262 | ||
|
|
bb5707f601 | ||
|
|
ff26836cab | ||
|
|
22037243ab | ||
|
|
acb662b5af | ||
|
|
e926f7a069 | ||
|
|
e01e34da1f | ||
|
|
47514f678d | ||
|
|
907c9e1cdd | ||
|
|
de3c9c895a | ||
|
|
4576ddb35e | ||
|
|
68392684b5 | ||
|
|
e4a4d22537 | ||
|
|
a3b2f67337 | ||
|
|
1e00c8fb28 | ||
|
|
ff16dce5c0 | ||
|
|
f972b31bf2 | ||
|
|
3dacd6a719 | ||
|
|
7bf99791ad | ||
|
|
7a5587b5af | ||
|
|
d6cf172846 | ||
|
|
cf86f4f0a4 | ||
|
|
b1f8309a20 | ||
|
|
8992c6af64 | ||
|
|
e4084a961b | ||
|
|
3ec1edefbe | ||
|
|
d3f33a9f09 | ||
|
|
d0ef3b5601 |
+13
-3
@@ -8,6 +8,8 @@ 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()
|
||||
@@ -216,9 +218,11 @@ 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", )
|
||||
@@ -227,10 +231,16 @@ 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_negative_text_embeds, ovi_audio_cfg):
|
||||
negative_text_embeds = ovi_negative_text_embeds.get("negative_prompt_embeds", None)
|
||||
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)
|
||||
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({
|
||||
|
||||
+1437
-1869
File diff suppressed because it is too large
Load Diff
+1507
File diff suppressed because it is too large
Load Diff
@@ -22,30 +22,6 @@ 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
|
||||
@@ -789,6 +765,98 @@ 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
|
||||
@@ -859,10 +927,6 @@ 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
|
||||
@@ -980,7 +1044,7 @@ class WanVideoImageToVideoEncode:
|
||||
gc.collect()
|
||||
|
||||
image_embeds = {
|
||||
"image_embeds": y,
|
||||
"image_embeds": y.cpu(),
|
||||
"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,
|
||||
@@ -992,7 +1056,7 @@ class WanVideoImageToVideoEncode:
|
||||
"fun_or_fl2v_model": fun_or_fl2v_model,
|
||||
"has_ref": has_ref,
|
||||
"add_cond_latents": add_cond_latents,
|
||||
"mask": mask
|
||||
"mask": mask.cpu()
|
||||
}
|
||||
|
||||
return (image_embeds,)
|
||||
@@ -1182,6 +1246,46 @@ 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
|
||||
@@ -2230,7 +2334,9 @@ NODE_CLASS_MAPPINGS = {
|
||||
"WanVideoAnimateEmbeds": WanVideoAnimateEmbeds,
|
||||
"WanVideoAddLucyEditLatents": WanVideoAddLucyEditLatents,
|
||||
"WanVideoSchedulerSA_ODE": WanVideoSchedulerSA_ODE,
|
||||
"WanVideoAddVideoPromptEmbeds": WanVideoAddVideoPromptEmbeds,
|
||||
"WanVideoAddBindweaveEmbeds": WanVideoAddBindweaveEmbeds,
|
||||
"TextImageEncodeQwenVL": TextImageEncodeQwenVL,
|
||||
"WanVideoUniLumosEmbeds": WanVideoUniLumosEmbeds,
|
||||
}
|
||||
|
||||
NODE_DISPLAY_NAME_MAPPINGS = {
|
||||
@@ -2270,4 +2376,6 @@ 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",
|
||||
}
|
||||
|
||||
+23
-11
@@ -395,7 +395,7 @@ class WanVideoLoraSelect:
|
||||
return (loras_list,)
|
||||
|
||||
try:
|
||||
lora_path = folder_paths.get_full_path("loras", lora)
|
||||
lora_path = folder_paths.get_full_path_or_raise("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("loras", lora_name),
|
||||
"path": folder_paths.get_full_path_or_raise("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("diffusion_models", vace_model)}]
|
||||
vace_model = [{"path": folder_paths.get_full_path_or_raise("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("diffusion_models", extra_model)}
|
||||
extra_model = {"path": folder_paths.get_full_path_or_raise("diffusion_models", extra_model)}
|
||||
if prev_model is not None and isinstance(prev_model, list):
|
||||
extra_model_list = prev_model + [extra_model]
|
||||
else:
|
||||
@@ -1088,6 +1088,14 @@ 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:
|
||||
@@ -1140,7 +1148,6 @@ 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"])
|
||||
@@ -1349,7 +1356,6 @@ 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
|
||||
|
||||
}
|
||||
|
||||
@@ -1480,6 +1486,12 @@ 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),
|
||||
@@ -1677,7 +1689,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("vae", model_name)
|
||||
model_path = folder_paths.get_full_path_or_raise("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())
|
||||
@@ -1728,7 +1740,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("vae_approx", model_name)
|
||||
model_path = folder_paths.get_full_path_or_raise("vae_approx", model_name)
|
||||
vae_sd = load_torch_file(model_path, safe_load=True)
|
||||
|
||||
vae = TAEHV(vae_sd, parallel=parallel, dtype=dtype)
|
||||
@@ -1766,7 +1778,7 @@ class LoadWanVideoT5TextEncoder:
|
||||
|
||||
dtype = {"bf16": torch.bfloat16, "fp16": torch.float16, "fp32": torch.float32}[precision]
|
||||
|
||||
model_path = folder_paths.get_full_path("text_encoders", model_name)
|
||||
model_path = folder_paths.get_full_path_or_raise("text_encoders", model_name)
|
||||
sd = load_torch_file(model_path, safe_load=True)
|
||||
|
||||
if quantization == "disabled":
|
||||
@@ -1876,10 +1888,10 @@ class LoadWanVideoClipTextEncoder:
|
||||
|
||||
dtype = {"bf16": torch.bfloat16, "fp16": torch.float16, "fp32": torch.float32}[precision]
|
||||
|
||||
model_path = folder_paths.get_full_path("clip_vision", model_name)
|
||||
model_path = folder_paths.get_full_path_or_raise("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("text_encoders", model_name)
|
||||
model_path = folder_paths.get_full_path_or_raise("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")
|
||||
|
||||
+49
-39
@@ -296,7 +296,6 @@ 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)
|
||||
@@ -343,7 +342,8 @@ class WanVideoSampler:
|
||||
dtype=torch.float32,
|
||||
generator=seed_g,
|
||||
device=torch.device("cpu"))
|
||||
seq_len = image_embeds["max_seq_len"]
|
||||
|
||||
seq_len = math.ceil((noise.shape[2] * noise.shape[3]) / 4 * noise.shape[1])
|
||||
|
||||
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,22 +484,13 @@ 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)
|
||||
|
||||
@@ -925,6 +916,9 @@ 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()
|
||||
@@ -1050,7 +1044,7 @@ class WanVideoSampler:
|
||||
if standin_input is not None:
|
||||
rope_function = "comfy" # only works with this currently
|
||||
|
||||
freqs = freqs_mot_ref = None
|
||||
freqs = None
|
||||
transformer.rope_embedder.k = None
|
||||
transformer.rope_embedder.num_frames = None
|
||||
d = transformer.dim // transformer.num_heads
|
||||
@@ -1064,19 +1058,15 @@ 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).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)
|
||||
],
|
||||
dim=1)
|
||||
elif "comfy" in rope_function: # comfy's rope
|
||||
transformer.rope_embedder.k = riflex_freq_index
|
||||
transformer.rope_embedder.num_frames = latent_video_length
|
||||
@@ -1151,6 +1141,16 @@ 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,9 +1361,18 @@ class WanVideoSampler:
|
||||
z = z * c_in
|
||||
timestep = c_noise
|
||||
|
||||
x_mot_ref_input = None
|
||||
if image_cond_mot_ref is not None:
|
||||
x_mot_ref_input = [x_mot_ref.to(z)]
|
||||
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)
|
||||
|
||||
base_params = {
|
||||
'x': [z], # latent
|
||||
@@ -1420,11 +1429,7 @@ 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, # 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
|
||||
"num_cond_latents": len(all_indices) if transformer.is_longcat else None,
|
||||
}
|
||||
|
||||
batch_size = 1
|
||||
@@ -1440,6 +1445,7 @@ 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,
|
||||
@@ -1456,6 +1462,7 @@ 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:
|
||||
@@ -1579,23 +1586,23 @@ class WanVideoSampler:
|
||||
noise_pred_uncond.view(batch_size, -1)
|
||||
).view(batch_size, 1, 1, 1)
|
||||
|
||||
noise_pred_uncond = noise_pred_uncond * alpha
|
||||
noise_pred_uncond_scaled = noise_pred_uncond * alpha
|
||||
|
||||
if use_tangential:
|
||||
noise_pred_uncond = tangential_projection(noise_pred_cond, noise_pred_uncond)
|
||||
noise_pred_uncond_scaled = tangential_projection(noise_pred_cond, noise_pred_uncond_scaled)
|
||||
|
||||
# RAAG (RATIO-aware Adaptive Guidance)
|
||||
if raag_alpha > 0.0:
|
||||
cfg_scale = get_raag_guidance(noise_pred_cond, noise_pred_uncond, cfg_scale, raag_alpha)
|
||||
cfg_scale = get_raag_guidance(noise_pred_cond, noise_pred_uncond_scaled, 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 + cfg_scale * filtered_cond * alpha
|
||||
noise_pred = noise_pred_uncond_scaled + cfg_scale * filtered_cond * alpha
|
||||
else:
|
||||
noise_pred = noise_pred_uncond + cfg_scale * (noise_pred_cond - noise_pred_uncond)
|
||||
del noise_pred_uncond, noise_pred_cond
|
||||
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
|
||||
|
||||
if latent_model_input_ovi is not None:
|
||||
if ovi_audio_cfg is None:
|
||||
@@ -2984,6 +2991,9 @@ 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)
|
||||
|
||||
|
||||
@@ -1,6 +1,8 @@
|
||||
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
|
||||
|
||||
@@ -12,6 +14,9 @@ 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):
|
||||
@@ -660,6 +665,96 @@ 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,
|
||||
@@ -673,6 +768,7 @@ NODE_CLASS_MAPPINGS = {
|
||||
"NormalizeAudioLoudness": NormalizeAudioLoudness,
|
||||
"WanVideoPassImagesFromSamples": WanVideoPassImagesFromSamples,
|
||||
"FaceMaskFromPoseKeypoints": FaceMaskFromPoseKeypoints,
|
||||
"DrawGaussianNoiseOnImage": DrawGaussianNoiseOnImage,
|
||||
}
|
||||
NODE_DISPLAY_NAME_MAPPINGS = {
|
||||
"WanVideoImageResizeToClosest": "WanVideo Image Resize To Closest",
|
||||
@@ -686,4 +782,5 @@ 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
@@ -1,7 +1,7 @@
|
||||
[project]
|
||||
name = "ComfyUI-WanVideoWrapper"
|
||||
description = "ComfyUI wrapper nodes for WanVideo"
|
||||
version = "1.3.8"
|
||||
version = "1.3.9"
|
||||
license = {file = "LICENSE"}
|
||||
dependencies = ["accelerate >= 1.2.1", "diffusers >= 0.33.0", "peft >= 0.17.0", "ftfy", "gguf >= 0.17.1", "pyloudnorm"]
|
||||
|
||||
|
||||
@@ -1,7 +1,29 @@
|
||||
## 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.
|
||||
|
||||
## 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.
|
||||
|
||||
# WORK IN PROGRESS (perpetually)
|
||||
|
||||
# Why should I use custom nodes when WanVideo works natively?
|
||||
|
||||
+76
-183
@@ -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, mot_ref_latent=None):
|
||||
def rope_params(max_seq_len, dim, theta=10000, L_test=25, k=0, freqs_scaling=1.0):
|
||||
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,16 +259,9 @@ 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]
|
||||
|
||||
else:
|
||||
freqs = torch.outer(torch.arange(max_seq_len), inv_theta_pow)
|
||||
freqs = torch.polar(torch.ones_like(freqs), freqs)
|
||||
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)
|
||||
@@ -363,10 +356,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).to(x.dtype)
|
||||
return self._norm(x.to(self.weight.dtype)) * self.weight
|
||||
|
||||
def _norm(self, x):
|
||||
return x * (torch.rsqrt(x.pow(2).mean(dim=-1, keepdim=True) + self.eps))
|
||||
return x * (torch.rsqrt(x.pow(2).mean(dim=-1, keepdim=True) + self.eps)).to(x.dtype)
|
||||
|
||||
def forward_chunked(self, x, num_chunks=4):
|
||||
output = torch.empty_like(x)
|
||||
@@ -393,7 +386,7 @@ class WanFusedRMSNorm(nn.RMSNorm):
|
||||
if use_chunked:
|
||||
return self.forward_chunked(x, num_chunks)
|
||||
else:
|
||||
return super().forward(x.to(self.weight.dtype).to(x.dtype))
|
||||
return super().forward(x)
|
||||
|
||||
def forward_chunked(self, x, num_chunks=4):
|
||||
output = torch.empty_like(x)
|
||||
@@ -405,7 +398,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.to(self.weight.dtype)).to(chunk.dtype)
|
||||
output[:, start_idx:end_idx, :] = super().forward(chunk)
|
||||
start_idx = end_idx
|
||||
|
||||
return output
|
||||
@@ -471,8 +464,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)).view(b, s, n, d)
|
||||
k = self.norm_k(self.k(x)).view(b, s, n, d)
|
||||
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)
|
||||
v = self.v(x).view(b, s, n, d)
|
||||
return q, k, v
|
||||
|
||||
@@ -487,8 +480,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)).view(b, s, n, d)
|
||||
k = self.norm_k(self.k(x) + self.k_loras(x)).view(b, s, n, d)
|
||||
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)
|
||||
v = (self.v(x) + self.v_loras(x)).view(b, s, n, d)
|
||||
return q, k, v
|
||||
|
||||
@@ -681,7 +674,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).to(self.norm_q.weight.dtype)).to(x.dtype)
|
||||
q = self.norm_q(self.q(x).view(b, -1, n, d))
|
||||
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)
|
||||
|
||||
@@ -754,32 +747,6 @@ 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):
|
||||
|
||||
@@ -806,13 +773,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)).view(b, -1, n, d)
|
||||
k = self.norm_k(self.k(context).to(self.norm_k.weight.dtype)).view(b, -1, n, d).to(x.dtype)
|
||||
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)).view(b, -1, n, d)
|
||||
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)
|
||||
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
|
||||
@@ -904,8 +871,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).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)
|
||||
q = self.norm_q(self.q(x)).view(b, -1, n, d)
|
||||
k = self.norm_k(self.k(mo)).view(b, n, -1, d)
|
||||
v = self.v(mo).view(b, -1, n, d)
|
||||
|
||||
# compute attention
|
||||
@@ -930,7 +897,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, mot_ref_block=False, is_longcat=False):
|
||||
block_idx=0, is_longcat=False):
|
||||
super().__init__()
|
||||
self.dim = out_features
|
||||
self.ffn_dim = ffn_dim
|
||||
@@ -946,7 +913,6 @@ 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
|
||||
@@ -986,16 +952,6 @@ 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)
|
||||
@@ -1089,8 +1045,7 @@ 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]
|
||||
@@ -1126,12 +1081,6 @@ 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)
|
||||
@@ -1176,13 +1125,8 @@ 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":
|
||||
@@ -1192,9 +1136,6 @@ 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)
|
||||
@@ -1255,20 +1196,6 @@ 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)
|
||||
|
||||
@@ -1300,11 +1227,8 @@ 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
|
||||
@@ -1339,10 +1263,6 @@ 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,
|
||||
@@ -1350,8 +1270,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.to(self.norm4.weight.dtype)).to(input_dtype), 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), mtv_motion_tokens, mtv_motion_rotary_emb, grid_sizes, mtv_freqs)
|
||||
x = x.add(x_motion, alpha=mtv_strength)
|
||||
|
||||
# HuMo Audio Cross-Attention
|
||||
@@ -1397,13 +1317,7 @@ 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)
|
||||
|
||||
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
|
||||
return x, x_ip, lynx_ref_feature, x_ovi
|
||||
|
||||
@torch.compiler.disable()
|
||||
def split_cross_attn_ffn(self, x, context, shift_mlp, scale_mlp, gate_mlp, clip_embed=None, grid_sizes=None):
|
||||
@@ -1603,10 +1517,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, dtype=torch.float32):
|
||||
def forward(self, image_embeds):
|
||||
if hasattr(self, 'emb_pos'):
|
||||
image_embeds = image_embeds + self.emb_pos.to(image_embeds.device)
|
||||
clip_extra_context_tokens = self.proj(image_embeds.to(self.proj[1].weight.dtype)).to(dtype)
|
||||
clip_extra_context_tokens = self.proj(image_embeds)
|
||||
return clip_extra_context_tokens
|
||||
|
||||
from .s2v.auxi_blocks import MotionEncoder_tc
|
||||
@@ -1708,22 +1622,47 @@ 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,
|
||||
@@ -1731,8 +1670,8 @@ class WanModel(torch.nn.Module):
|
||||
# lynx
|
||||
lynx_ip_layers=None,
|
||||
lynx_ref_layers=None,
|
||||
# VAP
|
||||
is_VAP = False,
|
||||
# ovi
|
||||
is_ovi_audio_model=False,
|
||||
# LongCat
|
||||
is_longcat=False,
|
||||
):
|
||||
@@ -1869,7 +1808,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
|
||||
|
||||
@@ -1886,13 +1825,6 @@ 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
|
||||
@@ -1933,7 +1865,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, mot_ref_block=is_VAP and i in VAP_layers)
|
||||
block_idx=i, is_longcat=is_longcat)
|
||||
for i in range(num_layers)
|
||||
])
|
||||
#MTV Crafter
|
||||
@@ -2207,15 +2139,12 @@ 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, mot_ref=False):
|
||||
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):
|
||||
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:
|
||||
@@ -2297,7 +2226,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,
|
||||
x_mot_ref=None, mot_ref_context=None, mot_ref_clip_embeds=None, freqs_mot_ref=None,
|
||||
add_text_emb=None,
|
||||
):
|
||||
r"""
|
||||
Forward pass through the diffusion model
|
||||
@@ -2328,7 +2257,6 @@ 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
|
||||
@@ -2419,18 +2347,6 @@ 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
|
||||
@@ -2541,16 +2457,13 @@ 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
|
||||
@@ -2591,10 +2504,6 @@ 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]:
|
||||
@@ -2659,20 +2568,13 @@ class WanModel(torch.nn.Module):
|
||||
e = e.to(self.offload_device, non_blocking=self.use_non_blocking)
|
||||
|
||||
# clip vision embedding
|
||||
clip_embed = clip_embed_mot_ref = None
|
||||
clip_embed = 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
|
||||
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
|
||||
#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)
|
||||
|
||||
@@ -2698,13 +2600,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 = 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))
|
||||
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)
|
||||
|
||||
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 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 self.is_longcat:
|
||||
context[:, tokens:] = 0
|
||||
@@ -2948,13 +2850,6 @@ 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:
|
||||
@@ -3042,9 +2937,7 @@ 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
|
||||
# ====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=====
|
||||
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
|
||||
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:
|
||||
|
||||
Reference in New Issue
Block a user