diff --git a/nodes.py b/nodes.py index ca0ddf6..0a611e9 100644 --- a/nodes.py +++ b/nodes.py @@ -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,) @@ -2246,6 +2334,8 @@ NODE_CLASS_MAPPINGS = { "WanVideoAnimateEmbeds": WanVideoAnimateEmbeds, "WanVideoAddLucyEditLatents": WanVideoAddLucyEditLatents, "WanVideoSchedulerSA_ODE": WanVideoSchedulerSA_ODE, + "WanVideoAddBindweaveEmbeds": WanVideoAddBindweaveEmbeds, + "TextImageEncodeQwenVL": TextImageEncodeQwenVL, "WanVideoUniLumosEmbeds": WanVideoUniLumosEmbeds, } @@ -2286,5 +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", } diff --git a/nodes_model_loading.py b/nodes_model_loading.py index cd348b6..32a2ff4 100644 --- a/nodes_model_loading.py +++ b/nodes_model_loading.py @@ -1486,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), diff --git a/nodes_sampler.py b/nodes_sampler.py index 5a549e1..b8cd80a 100644 --- a/nodes_sampler.py +++ b/nodes_sampler.py @@ -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() @@ -1357,6 +1361,16 @@ 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) @@ -1415,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 @@ -1431,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, @@ -1447,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: @@ -2975,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) diff --git a/wanvideo/modules/model.py b/wanvideo/modules/model.py index 3358fc3..7a7c916 100644 --- a/wanvideo/modules/model.py +++ b/wanvideo/modules/model.py @@ -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