29 Commits
Author SHA1 Message Date
kijai f685ee33ac Merge branch 'main' into bindweave 2025-11-13 16:37:38 +02:00
kijai 4e31081262 Better errors when trying to load models that don't exist 2025-11-13 16:19:05 +02:00
kijai bb5707f601 Merge branch 'main' into bindweave 2025-11-11 18:53:19 +02:00
kijai ff26836cab Create wanvideo_2_2_5B_Ovi_image_to_video_audio_10_seconds_example_01.json 2025-11-11 18:53:11 +02:00
kijai 22037243ab Fix Ovi audio negative prompt
Had rather bad bug here which made Ovi audio always use the video negative prompt...
2025-11-11 17:47:33 +02:00
kijai acb662b5af Merge branch 'main' into bindweave 2025-11-11 11:44:26 +02:00
kijai e926f7a069 version bump 1.3.9 2025-11-11 10:57:46 +02:00
kijai e01e34da1f Update nodes_model_loading.py 2025-11-11 10:01:06 +02:00
kijai 47514f678d Allow loading original Ovi -models 2025-11-11 09:46:13 +02:00
kijai 907c9e1cdd Update nodes.py 2025-11-10 21:02:58 +02:00
kijai de3c9c895a Create wanvideo_1_3B_UniLumos_relight_example_01.json 2025-11-10 19:41:22 +02:00
kijai 4576ddb35e Add node to create input for UniLumos 2025-11-10 19:41:19 +02:00
kijai 68392684b5 Add node to use UniLumos
Simply allows fore and background latent inputs for UniLumos relight model, example inputs seem to work: https://github.com/alibaba-damo-academy/Lumos-Custom/tree/main/UniLumos/UniLumos/examples
2025-11-10 19:06:45 +02:00
kijai e4a4d22537 Update nodes_sampler.py 2025-11-08 16:04:55 +02:00
kijai a3b2f67337 Pad clip vision embeds like in original code 2025-11-08 16:03:00 +02:00
kijai 1e00c8fb28 Update nodes.py 2025-11-08 12:21:11 +02:00
kijai ff16dce5c0 Update nodes_sampler.py 2025-11-07 01:15:13 +02:00
kijai f972b31bf2 Update nodes.py 2025-11-07 01:09:22 +02:00
kijai 3dacd6a719 Update nodes_sampler.py 2025-11-07 00:45:06 +02:00
kijai 7bf99791ad Update nodes.py 2025-11-07 00:44:13 +02:00
kijai 7a5587b5af Let the user resize for QwenVL
Seems to need smaller resolutions
2025-11-07 00:39:10 +02:00
kijai d6cf172846 Update nodes.py 2025-11-06 23:41:24 +02:00
kijai cf86f4f0a4 Update model.py 2025-11-06 21:01:03 +02:00
kijai b1f8309a20 Update nodes_model_loading.py 2025-11-06 19:55:54 +02:00
kijai 8992c6af64 Don't include padding for scheduler 2025-11-06 19:22:56 +02:00
kijai e4084a961b Update nodes.py 2025-11-06 18:32:53 +02:00
kijai 3ec1edefbe init
For testing, no idea if it works yet
2025-11-06 17:35:48 +02:00
Jukka Seppänen d3f33a9f09 Update readme.md 2025-11-06 16:38:46 +02:00
kijai d0ef3b5601 Update readme.md 2025-11-06 16:37:50 +02:00
10 changed files with 3785 additions and 27 deletions
+13 -3
View File
@@ -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
+139 -6
View File
@@ -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
View File
@@ -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
View File
@@ -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)
+97
View File
@@ -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
View File
@@ -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"]
+24 -2
View File
@@ -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?
+8 -2
View File
@@ -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