fix control lora
This commit is contained in:
@@ -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
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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:
|
||||
|
||||
Reference in New Issue
Block a user