From 4aac86a828204c2bd1effb01ea137a798d4d9a46 Mon Sep 17 00:00:00 2001 From: kijai <40791699+kijai@users.noreply.github.com> Date: Tue, 19 Aug 2025 19:11:33 +0300 Subject: [PATCH] offloading fixes --- nodes.py | 130 ++++++++++++++++++++--------------------- nodes_model_loading.py | 5 +- 2 files changed, 66 insertions(+), 69 deletions(-) diff --git a/nodes.py b/nodes.py index 641fcde..55bdcb0 100644 --- a/nodes.py +++ b/nodes.py @@ -2218,38 +2218,38 @@ class WanVideoSampler: gc.collect() #blockswap init - if block_swap_args is not None and not transformer.patched_linear: - transformer.use_non_blocking = block_swap_args.get("use_non_blocking", False) - for name, param in transformer.named_parameters(): - if "block" not in name: - param.data = param.data.to(device) - if "control_adapter" in name: - param.data = param.data.to(device) - elif block_swap_args["offload_txt_emb"] and "txt_emb" in name: - param.data = param.data.to(offload_device) - elif block_swap_args["offload_img_emb"] and "img_emb" in name: - param.data = param.data.to(offload_device) + if not transformer.patched_linear: + if block_swap_args is not None: + transformer.use_non_blocking = block_swap_args.get("use_non_blocking", False) + for name, param in transformer.named_parameters(): + if "block" not in name: + param.data = param.data.to(device) + if "control_adapter" in name: + param.data = param.data.to(device) + elif block_swap_args["offload_txt_emb"] and "txt_emb" in name: + param.data = param.data.to(offload_device) + elif block_swap_args["offload_img_emb"] and "img_emb" in name: + param.data = param.data.to(offload_device) - transformer.block_swap( - block_swap_args["blocks_to_swap"] - 1 , - block_swap_args["offload_txt_emb"], - block_swap_args["offload_img_emb"], - vace_blocks_to_swap = block_swap_args.get("vace_blocks_to_swap", None), - prefetch_blocks = block_swap_args.get("prefetch_blocks", 0), - block_swap_debug = block_swap_args.get("block_swap_debug", False), - ) - elif model["auto_cpu_offload"]: - for module in transformer.modules(): - if hasattr(module, "offload"): - module.offload() - if hasattr(module, "onload"): - module.onload() - for block in transformer.blocks: - block.modulation = torch.nn.Parameter(block.modulation.to(device)) - transformer.head.modulation = torch.nn.Parameter(transformer.head.modulation.to(device)) - - #elif model["manual_offloading"]: - # transformer.to(device) + transformer.block_swap( + block_swap_args["blocks_to_swap"] - 1 , + block_swap_args["offload_txt_emb"], + block_swap_args["offload_img_emb"], + vace_blocks_to_swap = block_swap_args.get("vace_blocks_to_swap", None), + prefetch_blocks = block_swap_args.get("prefetch_blocks", 0), + block_swap_debug = block_swap_args.get("block_swap_debug", False), + ) + elif model["auto_cpu_offload"]: + for module in transformer.modules(): + if hasattr(module, "offload"): + module.offload() + if hasattr(module, "onload"): + module.onload() + for block in transformer.blocks: + block.modulation = torch.nn.Parameter(block.modulation.to(device)) + transformer.head.modulation = torch.nn.Parameter(transformer.head.modulation.to(device)) + else: + transformer.to(device) # Initialize Cache if enabled transformer.enable_teacache = transformer.enable_magcache = transformer.enable_easycache = False @@ -2657,7 +2657,7 @@ class WanVideoSampler: except Exception as e: log.error(f"Error during model prediction: {e}") if force_offload: - if model["manual_offloading"]: + if not model["auto_cpu_offload"]: offload_transformer(transformer) raise e @@ -3256,36 +3256,36 @@ class WanVideoSampler: if offload: #blockswap init - if transformer_options is not None: - block_swap_args = transformer_options.get("block_swap_args", None) + if not transformer.patched_linear: + if block_swap_args is not None: + transformer.use_non_blocking = block_swap_args.get("use_non_blocking", False) + for name, param in transformer.named_parameters(): + if "block" not in name: + param.data = param.data.to(device) + if "control_adapter" in name: + param.data = param.data.to(device) + elif block_swap_args["offload_txt_emb"] and "txt_emb" in name: + param.data = param.data.to(offload_device) + elif block_swap_args["offload_img_emb"] and "img_emb" in name: + param.data = param.data.to(offload_device) - if block_swap_args is not None: - transformer.use_non_blocking = block_swap_args.get("use_non_blocking", False) - for name, param in transformer.named_parameters(): - if "block" not in name: - param.data = param.data.to(device) - if "control_adapter" in name: - param.data = param.data.to(device) - elif block_swap_args["offload_txt_emb"] and "txt_emb" in name: - param.data = param.data.to(offload_device) - elif block_swap_args["offload_img_emb"] and "img_emb" in name: - param.data = param.data.to(offload_device) - - transformer.block_swap( - block_swap_args["blocks_to_swap"] - 1 , - block_swap_args["offload_txt_emb"], - block_swap_args["offload_img_emb"], - vace_blocks_to_swap = block_swap_args.get("vace_blocks_to_swap", None), - ) - - elif model["auto_cpu_offload"]: - for module in transformer.modules(): - if hasattr(module, "offload"): - module.offload() - if hasattr(module, "onload"): - module.onload() - elif model["manual_offloading"]: - transformer.to(device) + transformer.block_swap( + block_swap_args["blocks_to_swap"] - 1 , + block_swap_args["offload_txt_emb"], + block_swap_args["offload_img_emb"], + vace_blocks_to_swap = block_swap_args.get("vace_blocks_to_swap", None), + ) + elif model["auto_cpu_offload"]: + for module in transformer.modules(): + if hasattr(module, "offload"): + module.offload() + if hasattr(module, "onload"): + module.onload() + for block in transformer.blocks: + block.modulation = torch.nn.Parameter(block.modulation.to(device)) + transformer.head.modulation = torch.nn.Parameter(transformer.head.modulation.to(device)) + else: + transformer.to(device) comfy_pbar = ProgressBar(len(timesteps)-1) for i in tqdm(range(len(timesteps)-1)): @@ -3337,7 +3337,7 @@ class WanVideoSampler: comfy_pbar.update(1) if offload: - transformer.to(offload_device) + offload_transformer(transformer) vae.to(device) videos = vae.decode(x0.unsqueeze(0).to(vae.dtype), device=device, tiled=tiled_vae) vae.to(offload_device) @@ -3410,7 +3410,7 @@ class WanVideoSampler: del noise, latent if force_offload: - if model["manual_offloading"]: + if not model["auto_cpu_offload"]: transformer.to(offload_device) mm.soft_empty_cache() gc.collect() @@ -3508,7 +3508,7 @@ class WanVideoSampler: except Exception as e: log.error(f"Error during sampling: {e}") if force_offload: - if model["manual_offloading"]: + if not model["auto_cpu_offload"]: offload_transformer(transformer) raise e @@ -3516,7 +3516,7 @@ class WanVideoSampler: cache_report(transformer, cache_args) if force_offload: - if model["manual_offloading"]: + if not model["auto_cpu_offload"]: offload_transformer(transformer) try: diff --git a/nodes_model_loading.py b/nodes_model_loading.py index 60deda7..aaa5209 100644 --- a/nodes_model_loading.py +++ b/nodes_model_loading.py @@ -948,7 +948,7 @@ class WanVideoModelLoader: mm.unload_all_models() mm.cleanup_models() mm.soft_empty_cache() - manual_offloading = True + if "sage" in attention_mode: try: from sageattention import sageattn @@ -964,8 +964,6 @@ class WanVideoModelLoader: if merge_loras is True: raise ValueError("GGUF models do not support LoRA merging, please disable merge_loras in the LoRA select node.") - - manual_offloading = True transformer_load_device = device if load_device == "main_device" else offload_device base_dtype = {"fp8_e4m3fn": torch.float8_e4m3fn, "fp8_e4m3fn_fast": torch.float8_e4m3fn, "bf16": torch.bfloat16, "fp16": torch.float16, "fp16_fast": torch.float16, "fp32": torch.float32}[base_precision] @@ -1366,7 +1364,6 @@ class WanVideoModelLoader: patcher.model["weight_dtype"] = weight_dtype patcher.model["base_path"] = model_path patcher.model["model_name"] = model - patcher.model["manual_offloading"] = manual_offloading patcher.model["quantization"] = quantization patcher.model["auto_cpu_offload"] = True if vram_management_args is not None else False patcher.model["control_lora"] = control_lora