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({
|
||||
|
||||
File diff suppressed because it is too large
Load Diff
+1507
File diff suppressed because it is too large
Load Diff
@@ -765,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
|
||||
@@ -835,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
|
||||
@@ -956,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,
|
||||
@@ -968,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,)
|
||||
@@ -1158,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
|
||||
@@ -2206,6 +2334,9 @@ NODE_CLASS_MAPPINGS = {
|
||||
"WanVideoAnimateEmbeds": WanVideoAnimateEmbeds,
|
||||
"WanVideoAddLucyEditLatents": WanVideoAddLucyEditLatents,
|
||||
"WanVideoSchedulerSA_ODE": WanVideoSchedulerSA_ODE,
|
||||
"WanVideoAddBindweaveEmbeds": WanVideoAddBindweaveEmbeds,
|
||||
"TextImageEncodeQwenVL": TextImageEncodeQwenVL,
|
||||
"WanVideoUniLumosEmbeds": WanVideoUniLumosEmbeds,
|
||||
}
|
||||
|
||||
NODE_DISPLAY_NAME_MAPPINGS = {
|
||||
@@ -2245,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
-10
@@ -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"])
|
||||
@@ -1479,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),
|
||||
@@ -1676,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())
|
||||
@@ -1727,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)
|
||||
@@ -1765,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":
|
||||
@@ -1875,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")
|
||||
|
||||
+35
-3
@@ -342,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:
|
||||
@@ -411,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)
|
||||
@@ -915,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()
|
||||
@@ -1137,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,
|
||||
@@ -1347,6 +1361,19 @@ 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)
|
||||
|
||||
base_params = {
|
||||
'x': [z], # latent
|
||||
'y': [image_cond_input] if image_cond_input is not None else None, # image cond
|
||||
@@ -1402,7 +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
|
||||
"num_cond_latents": len(all_indices) if transformer.is_longcat else None,
|
||||
}
|
||||
|
||||
batch_size = 1
|
||||
@@ -1418,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,
|
||||
@@ -1434,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:
|
||||
@@ -2962,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?
|
||||
|
||||
@@ -2226,6 +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,
|
||||
add_text_emb=None,
|
||||
):
|
||||
r"""
|
||||
Forward pass through the diffusion model
|
||||
@@ -2599,8 +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 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
|
||||
|
||||
Reference in New Issue
Block a user