diff --git a/nodes.py b/nodes.py index 8900ca8..6d20a44 100644 --- a/nodes.py +++ b/nodes.py @@ -1415,7 +1415,7 @@ class WanVideoSampler: image_cond = control_latents.to(device) if not patcher.model.is_patched: log.info("Re-loading control LoRA...") - patcher = apply_lora(patcher, device, device, low_mem_load=False) + patcher = apply_lora(patcher, device, device, low_mem_load=False, control_lora=True) patcher.model.is_patched = True else: if transformer.in_dim not in [48, 32]: @@ -1887,7 +1887,7 @@ class WanVideoSampler: image_cond_input = control_latents.to(z) if not patcher.model.is_patched: log.info("Loading LoRA...") - patcher = apply_lora(patcher, device, device, low_mem_load=False) + patcher = apply_lora(patcher, device, device, low_mem_load=False, control_lora=True) patcher.model.is_patched = True elif ATI_tracks is not None and ((ati_start_percent <= current_step_percentage <= ati_end_percent) or diff --git a/nodes_model_loading.py b/nodes_model_loading.py index 0febff0..6f89800 100644 --- a/nodes_model_loading.py +++ b/nodes_model_loading.py @@ -1080,7 +1080,7 @@ class WanVideoModelLoader: if not gguf and not "scaled" in quantization and merge_loras: log.info("Patching LoRA to the model...") - patcher = apply_lora(patcher, device, transformer_load_device, params_to_keep=params_to_keep, dtype=dtype, base_dtype=base_dtype, state_dict=sd, low_mem_load=lora_low_mem_load) + patcher = apply_lora(patcher, device, transformer_load_device, params_to_keep=params_to_keep, dtype=dtype, base_dtype=base_dtype, state_dict=sd, low_mem_load=lora_low_mem_load, control_lora=control_lora) if gguf: #from diffusers.quantizers.gguf.utils import _replace_with_gguf_linear, GGUFParameter diff --git a/utils.py b/utils.py index a24dd22..9a1c13e 100644 --- a/utils.py +++ b/utils.py @@ -2,7 +2,7 @@ import importlib.metadata import torch import logging from tqdm import tqdm -import types +import types, collections from comfy.utils import ProgressBar, copy_to_param, set_attr_param from comfy.model_patcher import get_key_weight, string_to_seed from comfy.lora import calculate_weight @@ -42,13 +42,16 @@ def get_tensor_memory(tensor): memory_bytes = tensor.element_size() * tensor.nelement() return f"{memory_bytes / (1024 * 1024):.2f} MB" -def patch_weight_to_device(self, key, device_to=None, inplace_update=False): +def patch_weight_to_device(self, key, device_to=None, inplace_update=False, backup_keys=False): if key not in self.patches: return - + weight, set_func, convert_func = get_key_weight(self.model, key) inplace_update = self.weight_inplace_update or inplace_update + if backup_keys and key not in self.backup: + self.backup[key] = collections.namedtuple('Dimension', ['weight', 'inplace_update'])(weight.to(device=self.offload_device, copy=inplace_update), inplace_update) + if device_to is not None: temp_weight = cast_to_device(weight, device_to, torch.float32, copy=True) else: @@ -66,7 +69,7 @@ def patch_weight_to_device(self, key, device_to=None, inplace_update=False): else: set_func(out_weight, inplace_update=inplace_update, seed=string_to_seed(key)) -def apply_lora(model, device_to, transformer_load_device, params_to_keep=None, dtype=None, base_dtype=None, state_dict=None, low_mem_load=False): +def apply_lora(model, device_to, transformer_load_device, params_to_keep=None, dtype=None, base_dtype=None, state_dict=None, low_mem_load=False, control_lora=False): model.patch_weight_to_device = types.MethodType(patch_weight_to_device, model) to_load = [] for n, m in model.model.named_modules(): @@ -104,10 +107,9 @@ def apply_lora(model, device_to, transformer_load_device, params_to_keep=None, d except: continue if low_mem_load: - model.patch_weight_to_device("{}.{}".format(name, param), device_to=device_to, inplace_update=True) + model.patch_weight_to_device("{}.{}".format(name, param), device_to=device_to, inplace_update=True, backup_keys=control_lora) else: - model.patch_weight_to_device("{}.{}".format(name, param), device_to=device_to) - model.backup["{}.{}".format(name, param)] = None + model.patch_weight_to_device("{}.{}".format(name, param), device_to=device_to, backup_keys=control_lora) if device_to != transformer_load_device: set_module_tensor_to_device(m, param, device=transformer_load_device) if low_mem_load: