diff --git a/nodes.py b/nodes.py index 89d8646..1afed65 100644 --- a/nodes.py +++ b/nodes.py @@ -775,7 +775,8 @@ class WanVideoAddBindweaveEmbeds: }, "optional": { "ref_masks": ("MASK", {"tooltip": "Reference mask to encode"}), - "qwenvl_embeds": ("QWENVL_EMBEDS", {"tooltip": "Qwen-VL image embeddings for the reference image"}), + "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"}), } } @@ -784,7 +785,7 @@ class WanVideoAddBindweaveEmbeds: FUNCTION = "add" CATEGORY = "WanVideoWrapper" - def add(self, embeds, reference_latents, ref_masks=None, qwenvl_embeds=None): + 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 @@ -809,7 +810,8 @@ class WanVideoAddBindweaveEmbeds: updated["mask"] = mask updated["image_embeds"] = image_embeds - updated["qwenvl_embeds"] = qwenvl_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(): @@ -818,6 +820,7 @@ class TextImageEncodeQwenVL(): return {"required": { "clip": ("CLIP",), "prompt": ("STRING", {"default": "", "multiline": True}), + "auto_resize": ("BOOLEAN", {"default": True, "tooltip": "Use the original code's image resize logic"}), }, "optional": { "image": ("IMAGE", ), @@ -829,22 +832,28 @@ class TextImageEncodeQwenVL(): FUNCTION = "add" CATEGORY = "WanVideoWrapper" - def add(cls, clip, prompt, image=None): + def add(cls, clip, prompt, auto_resize, image=None): if image is None: - images = [] + input_images = [] else: - samples = image.movedim(-1, 1) - total = int(1280 * 720) + if auto_resize: + total = int(1280 * 720) + width = image.shape[2] + height = image.shape[1] - 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) + if width * height > total: + new_width = int(width * 0.3) + new_height = int(height * 0.3) + else: + new_width = int(width * 0.75) + new_height = int(height * 0.75) - s = common_upscale(samples, width, height, "area", "disabled") - image = s.movedim(1, -1) - images = [image[:, :, :, :3]] + input_images = common_upscale(image.movedim(-1, 1), new_width, new_height, "area", "disabled").movedim(1, -1) + log.info(f"TextImageEncodeQwenVL: auto-resized image to: {new_width}x{new_height}") + else: + input_images = image - tokens = clip.tokenize(prompt, images=images) + tokens = clip.tokenize(prompt, images=input_images[..., :3]) conditioning = clip.encode_from_tokens_scheduled(tokens) print("Qwen-VL embeds shape:", conditioning[0][0].shape) diff --git a/nodes_sampler.py b/nodes_sampler.py index 2ff4c2c..a16aaa9 100644 --- a/nodes_sampler.py +++ b/nodes_sampler.py @@ -915,7 +915,8 @@ class WanVideoSampler: log.info(f"Number of shots in prompt: {shot_num}, Shot token lengths: {shot_len}") # Bindweave - qwenvl_embeds = image_embeds.get("qwenvl_embeds", None) + 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() @@ -1410,7 +1411,6 @@ class WanVideoSampler: "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, - "add_text_emb": qwenvl_embeds.to(device) if qwenvl_embeds is not None else None # QwenVL embeddings for Bindweave } batch_size = 1 @@ -1426,6 +1426,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 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, @@ -1442,6 +1443,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: