Compare commits
18
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
f685ee33ac | ||
|
|
bb5707f601 | ||
|
|
acb662b5af | ||
|
|
907c9e1cdd | ||
|
|
e4a4d22537 | ||
|
|
a3b2f67337 | ||
|
|
1e00c8fb28 | ||
|
|
ff16dce5c0 | ||
|
|
f972b31bf2 | ||
|
|
3dacd6a719 | ||
|
|
7bf99791ad | ||
|
|
7a5587b5af | ||
|
|
d6cf172846 | ||
|
|
cf86f4f0a4 | ||
|
|
b1f8309a20 | ||
|
|
8992c6af64 | ||
|
|
e4084a961b | ||
|
|
3ec1edefbe |
@@ -765,6 +765,98 @@ class WanVideoAddStandInLatent:
|
|||||||
updated = dict(embeds)
|
updated = dict(embeds)
|
||||||
updated["standin_input"] = new_entry
|
updated["standin_input"] = new_entry
|
||||||
return (updated,)
|
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:
|
class WanVideoAddMTVMotion:
|
||||||
@classmethod
|
@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,
|
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):
|
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:
|
if vae is None:
|
||||||
raise ValueError("VAE is required for image encoding.")
|
raise ValueError("VAE is required for image encoding.")
|
||||||
H = height
|
H = height
|
||||||
@@ -956,7 +1044,7 @@ class WanVideoImageToVideoEncode:
|
|||||||
gc.collect()
|
gc.collect()
|
||||||
|
|
||||||
image_embeds = {
|
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,
|
"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,
|
"negative_clip_context": clip_embeds.get("negative_clip_embeds", None) if clip_embeds is not None else None,
|
||||||
"max_seq_len": max_seq_len,
|
"max_seq_len": max_seq_len,
|
||||||
@@ -968,7 +1056,7 @@ class WanVideoImageToVideoEncode:
|
|||||||
"fun_or_fl2v_model": fun_or_fl2v_model,
|
"fun_or_fl2v_model": fun_or_fl2v_model,
|
||||||
"has_ref": has_ref,
|
"has_ref": has_ref,
|
||||||
"add_cond_latents": add_cond_latents,
|
"add_cond_latents": add_cond_latents,
|
||||||
"mask": mask
|
"mask": mask.cpu()
|
||||||
}
|
}
|
||||||
|
|
||||||
return (image_embeds,)
|
return (image_embeds,)
|
||||||
@@ -2246,6 +2334,8 @@ NODE_CLASS_MAPPINGS = {
|
|||||||
"WanVideoAnimateEmbeds": WanVideoAnimateEmbeds,
|
"WanVideoAnimateEmbeds": WanVideoAnimateEmbeds,
|
||||||
"WanVideoAddLucyEditLatents": WanVideoAddLucyEditLatents,
|
"WanVideoAddLucyEditLatents": WanVideoAddLucyEditLatents,
|
||||||
"WanVideoSchedulerSA_ODE": WanVideoSchedulerSA_ODE,
|
"WanVideoSchedulerSA_ODE": WanVideoSchedulerSA_ODE,
|
||||||
|
"WanVideoAddBindweaveEmbeds": WanVideoAddBindweaveEmbeds,
|
||||||
|
"TextImageEncodeQwenVL": TextImageEncodeQwenVL,
|
||||||
"WanVideoUniLumosEmbeds": WanVideoUniLumosEmbeds,
|
"WanVideoUniLumosEmbeds": WanVideoUniLumosEmbeds,
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -2286,5 +2376,6 @@ NODE_DISPLAY_NAME_MAPPINGS = {
|
|||||||
"WanVideoAnimateEmbeds": "WanVideo Animate Embeds",
|
"WanVideoAnimateEmbeds": "WanVideo Animate Embeds",
|
||||||
"WanVideoAddLucyEditLatents": "WanVideo Add LucyEdit Latents",
|
"WanVideoAddLucyEditLatents": "WanVideo Add LucyEdit Latents",
|
||||||
"WanVideoSchedulerSA_ODE": "WanVideo Scheduler SA-ODE",
|
"WanVideoSchedulerSA_ODE": "WanVideo Scheduler SA-ODE",
|
||||||
|
"WanVideoAddBindweaveEmbeds": "WanVideo Add Bindweave Embeds",
|
||||||
"WanVideoUniLumosEmbeds": "WanVideo UniLumos Embeds",
|
"WanVideoUniLumosEmbeds": "WanVideo UniLumos Embeds",
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -1486,6 +1486,12 @@ class WanVideoModelLoader:
|
|||||||
transformer.add_proj = zero_module(torch.nn.Linear(inner_dim, inner_dim))
|
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)
|
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
|
latent_format=Wan22 if dim == 3072 else Wan21
|
||||||
comfy_model = WanVideoModel(
|
comfy_model = WanVideoModel(
|
||||||
WanVideoModelConfig(base_dtype, latent_format=latent_format),
|
WanVideoModelConfig(base_dtype, latent_format=latent_format),
|
||||||
|
|||||||
+22
-3
@@ -342,7 +342,8 @@ class WanVideoSampler:
|
|||||||
dtype=torch.float32,
|
dtype=torch.float32,
|
||||||
generator=seed_g,
|
generator=seed_g,
|
||||||
device=torch.device("cpu"))
|
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)
|
control_embeds = image_embeds.get("control_embeds", None)
|
||||||
if control_embeds is not None:
|
if control_embeds is not None:
|
||||||
@@ -411,7 +412,7 @@ class WanVideoSampler:
|
|||||||
dtype=torch.float32,
|
dtype=torch.float32,
|
||||||
device=torch.device("cpu"),
|
device=torch.device("cpu"),
|
||||||
generator=seed_g)
|
generator=seed_g)
|
||||||
|
|
||||||
seq_len = math.ceil((noise.shape[2] * noise.shape[3]) / 4 * noise.shape[1])
|
seq_len = math.ceil((noise.shape[2] * noise.shape[3]) / 4 * noise.shape[1])
|
||||||
|
|
||||||
recammaster = image_embeds.get("recammaster", None)
|
recammaster = image_embeds.get("recammaster", None)
|
||||||
@@ -915,6 +916,9 @@ class WanVideoSampler:
|
|||||||
rope_function = "default" #echoshot does not support comfy rope function
|
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}")
|
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.unload_all_models()
|
||||||
mm.soft_empty_cache()
|
mm.soft_empty_cache()
|
||||||
@@ -1357,6 +1361,16 @@ class WanVideoSampler:
|
|||||||
z = z * c_in
|
z = z * c_in
|
||||||
timestep = c_noise
|
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:
|
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)
|
z = torch.cat([z, foreground_latents.to(z), background_latents.to(z)], dim=0)
|
||||||
|
|
||||||
@@ -1415,7 +1429,7 @@ class WanVideoSampler:
|
|||||||
"ovi_negative_text_embeds": ovi_negative_text_embeds, # Audio latent model negative text embeds for Ovi
|
"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_LQ_latent": flashvsr_LQ_latent, # FlashVSR LQ latent for upsampling
|
||||||
"flashvsr_strength": flashvsr_strength, # FlashVSR strength
|
"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
|
batch_size = 1
|
||||||
@@ -1431,6 +1445,7 @@ class WanVideoSampler:
|
|||||||
#conditional (positive) pass
|
#conditional (positive) pass
|
||||||
if pos_latent is not None: # for humo
|
if pos_latent is not None: # for humo
|
||||||
base_params['x'] = [torch.cat([z[:, :-humo_reference_count], pos_latent], dim=1)]
|
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(
|
noise_pred_cond, noise_pred_ovi, cache_state_cond = transformer(
|
||||||
context=positive_embeds,
|
context=positive_embeds,
|
||||||
pred_id=cache_state[0] if cache_state else None,
|
pred_id=cache_state[0] if cache_state else None,
|
||||||
@@ -1447,6 +1462,7 @@ class WanVideoSampler:
|
|||||||
#unconditional (negative) pass
|
#unconditional (negative) pass
|
||||||
base_params['is_uncond'] = True
|
base_params['is_uncond'] = True
|
||||||
base_params['clip_fea'] = clip_fea_neg if clip_fea_neg is not None else clip_fea
|
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:
|
if wananim_face_pixels is not None:
|
||||||
base_params['wananim_face_pixel_values'] = torch.zeros_like(wananim_face_pixels).to(device, torch.float32) - 1
|
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:
|
if humo_audio_input_neg is not None:
|
||||||
@@ -2975,6 +2991,9 @@ class WanVideoSampler:
|
|||||||
if flowedit_args is None:
|
if flowedit_args is None:
|
||||||
latent = latent.to(intermediate_device)
|
latent = latent.to(intermediate_device)
|
||||||
|
|
||||||
|
if self.noise_front_pad_num > 0:
|
||||||
|
noise_pred = noise_pred[:, self.noise_front_pad_num:]
|
||||||
|
|
||||||
if use_tsr:
|
if use_tsr:
|
||||||
noise_pred = temporal_score_rescaling(noise_pred, latent, timestep, tsr_k, tsr_sigma)
|
noise_pred = temporal_score_rescaling(noise_pred, latent, timestep, tsr_k, tsr_sigma)
|
||||||
|
|
||||||
|
|||||||
@@ -2226,6 +2226,7 @@ class WanModel(torch.nn.Module):
|
|||||||
x_ovi=None, seq_len_ovi=None, ovi_negative_text_embeds=None,
|
x_ovi=None, seq_len_ovi=None, ovi_negative_text_embeds=None,
|
||||||
flashvsr_LQ_latent=None, flashvsr_strength=1.0,
|
flashvsr_LQ_latent=None, flashvsr_strength=1.0,
|
||||||
num_cond_latents=None,
|
num_cond_latents=None,
|
||||||
|
add_text_emb=None,
|
||||||
):
|
):
|
||||||
r"""
|
r"""
|
||||||
Forward pass through the diffusion model
|
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))
|
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]
|
tokens = context[0].shape[0]
|
||||||
context = self.text_embedding(
|
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)
|
||||||
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:
|
if self.is_longcat:
|
||||||
context[:, tokens:] = 0
|
context[:, tokens:] = 0
|
||||||
|
|||||||
Reference in New Issue
Block a user