init
For testing, no idea if it works yet
This commit is contained in:
@@ -765,6 +765,90 @@ 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": ("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=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
|
||||
updated["image_embeds"] = image_embeds
|
||||
updated["qwenvl_embeds"] = qwenvl_embeds
|
||||
return (updated, {"samples": image_embeds.unsqueeze(0)}, mask[0])
|
||||
|
||||
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:
|
||||
images = []
|
||||
else:
|
||||
samples = image.movedim(-1, 1)
|
||||
total = int(1024 * 1024)
|
||||
|
||||
scale_by = math.sqrt(total / (samples.shape[3] * samples.shape[2]))
|
||||
width = round(samples.shape[3] * scale_by)
|
||||
height = round(samples.shape[2] * scale_by)
|
||||
|
||||
s = common_upscale(samples, width, height, "area", "disabled")
|
||||
image = s.movedim(1, -1)
|
||||
images = [image[:, :, :, :3]]
|
||||
|
||||
tokens = clip.tokenize(prompt, images=images)
|
||||
conditioning = clip.encode_from_tokens_scheduled(tokens)
|
||||
print("Qwen-VL embeds shape:", conditioning[0][0].shape)
|
||||
|
||||
return conditioning[0][0],
|
||||
|
||||
class WanVideoAddMTVMotion:
|
||||
@classmethod
|
||||
@@ -956,7 +1040,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 +1052,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,)
|
||||
@@ -2206,6 +2290,8 @@ NODE_CLASS_MAPPINGS = {
|
||||
"WanVideoAnimateEmbeds": WanVideoAnimateEmbeds,
|
||||
"WanVideoAddLucyEditLatents": WanVideoAddLucyEditLatents,
|
||||
"WanVideoSchedulerSA_ODE": WanVideoSchedulerSA_ODE,
|
||||
"WanVideoAddBindweaveEmbeds": WanVideoAddBindweaveEmbeds,
|
||||
"TextImageEncodeQwenVL": TextImageEncodeQwenVL,
|
||||
}
|
||||
|
||||
NODE_DISPLAY_NAME_MAPPINGS = {
|
||||
@@ -2245,4 +2331,5 @@ NODE_DISPLAY_NAME_MAPPINGS = {
|
||||
"WanVideoAnimateEmbeds": "WanVideo Animate Embeds",
|
||||
"WanVideoAddLucyEditLatents": "WanVideo Add LucyEdit Latents",
|
||||
"WanVideoSchedulerSA_ODE": "WanVideo Scheduler SA-ODE",
|
||||
"WanVideoAddBindweaveEmbeds": "WanVideo Add Bindweave Embeds",
|
||||
}
|
||||
|
||||
@@ -1479,6 +1479,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_projector 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),
|
||||
|
||||
+14
-4
@@ -342,7 +342,12 @@ class WanVideoSampler:
|
||||
dtype=torch.float32,
|
||||
generator=seed_g,
|
||||
device=torch.device("cpu"))
|
||||
seq_len = image_embeds["max_seq_len"]
|
||||
|
||||
noise_front_pad_num = image_cond.shape[1] - noise.shape[1]
|
||||
if noise_front_pad_num > 0:
|
||||
pad = torch.zeros((noise.shape[0], noise_front_pad_num, noise.shape[2], noise.shape[3]), dtype=noise.dtype, device=noise.device)
|
||||
noise = torch.concat([pad, noise], dim=1)
|
||||
|
||||
|
||||
control_embeds = image_embeds.get("control_embeds", None)
|
||||
if control_embeds is not None:
|
||||
@@ -411,8 +416,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)
|
||||
if recammaster is not None:
|
||||
@@ -863,6 +867,7 @@ class WanVideoSampler:
|
||||
seq_len = math.ceil((noise.shape[2] * noise.shape[3]) / 4 * noise.shape[1])
|
||||
|
||||
latent = noise
|
||||
seq_len = math.ceil((noise.shape[2] * noise.shape[3]) / 4 * noise.shape[1])
|
||||
|
||||
#controlnet
|
||||
controlnet_latents = controlnet = None
|
||||
@@ -915,6 +920,8 @@ 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 = image_embeds.get("qwenvl_embeds", None)
|
||||
|
||||
mm.unload_all_models()
|
||||
mm.soft_empty_cache()
|
||||
@@ -1402,7 +1409,8 @@ 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,
|
||||
"add_text_emb": qwenvl_embeds.to(device) if qwenvl_embeds is not None else None # QwenVL embeddings for Bindweave
|
||||
}
|
||||
|
||||
batch_size = 1
|
||||
@@ -3061,6 +3069,8 @@ class WanVideoSampler:
|
||||
latent = latent[:,:-phantom_latents.shape[1]]
|
||||
if humo_reference_count > 0:
|
||||
latent = latent[:,:-humo_reference_count]
|
||||
if noise_front_pad_num > 0:
|
||||
latent = latent[:, noise_front_pad_num:]
|
||||
|
||||
cache_states = None
|
||||
if cache_args is not None:
|
||||
|
||||
@@ -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([context, add_text_emb], dim=1)
|
||||
context = self.text_embedding(context)
|
||||
|
||||
if self.is_longcat:
|
||||
context[:, tokens:] = 0
|
||||
|
||||
Reference in New Issue
Block a user