From a4b1a4be7ac69ba705cd9cc43450751cd2e6f25c Mon Sep 17 00:00:00 2001 From: kijai <40791699+kijai@users.noreply.github.com> Date: Tue, 11 Mar 2025 09:57:46 +0200 Subject: [PATCH] Add optional offloading of the main model before text encoding --- nodes.py | 19 +++++++++++++++---- 1 file changed, 15 insertions(+), 4 deletions(-) diff --git a/nodes.py b/nodes.py index c1ef9c2..ded96d9 100644 --- a/nodes.py +++ b/nodes.py @@ -712,7 +712,7 @@ class LoadWanVideoT5TextEncoder: }, "optional": { "load_device": (["main_device", "offload_device"], {"default": "offload_device"}), - "quantization": (['disabled', 'fp8_e4m3fn'], {"default": 'disabled', "tooltip": "optional quantization method"}), + "quantization": (['disabled', 'fp8_e4m3fn'], {"default": 'disabled', "tooltip": "optional quantization method"}), } } @@ -800,6 +800,7 @@ class WanVideoTextEncode: }, "optional": { "force_offload": ("BOOLEAN", {"default": True}), + "model_to_offload": ("WANVIDEOMODEL", {"tooltip": "Model to move to offload_device before encoding"}), } } @@ -809,10 +810,16 @@ class WanVideoTextEncode: CATEGORY = "WanVideoWrapper" DESCRIPTION = "Encodes text prompts into text embeddings. For context windowing you can input multiple prompts separated by '|'" - def process(self, t5, positive_prompt, negative_prompt,force_offload=True): + def process(self, t5, positive_prompt, negative_prompt,force_offload=True, model_to_offload=None): device = mm.get_torch_device() offload_device = mm.unet_offload_device() + + if model_to_offload is not None: + log.info(f"Moving video model to {offload_device}") + model_to_offload.model.to(offload_device) + mm.soft_empty_cache() + encoder = t5["model"] dtype = t5["dtype"] @@ -1682,8 +1689,12 @@ class WanVideoSampler: if image_cond is not None: partial_img_emb = image_cond[:, c, :, :] partial_image_cond = image_cond[:, 0, :, :].to(intermediate_device) - if min(c) > len(c) // 2: - partial_image_cond *= 0.1 + if min(c) > 0: #wip + control_strength = 1.0 + fade_rate = 0.01 + frame_position = min(c) + strength = max(control_strength * (1.0 - (frame_position * fade_rate)), 0.1) + partial_image_cond *= strength partial_img_emb[:, 0, :, :] = partial_image_cond partial_latent_model_input = latent_model_input[:, c, :, :]