simplify LoRA loading further

This commit is contained in:
kijai
2025-03-12 19:46:01 +02:00
parent 0c780581f0
commit b505075850
2 changed files with 40 additions and 6 deletions
+7 -5
View File
@@ -1,7 +1,7 @@
import os
import torch
import gc
from .utils import log, print_memory
from .utils import log, print_memory, apply_lora
import numpy as np
import math
from tqdm import tqdm
@@ -424,6 +424,8 @@ class WanVideoModelLoader:
else:
dtype = base_dtype
params_to_keep = {"norm", "head", "bias", "time_in", "vector_in", "patch_embedding", "time_", "img_emb", "modulation"}
if lora is not None:
transformer_load_device = device
for name, param in transformer.named_parameters():
#print("Assigning Parameter name: ", name)
dtype_to_use = base_dtype if any(keyword in name for keyword in params_to_keep) else dtype
@@ -477,9 +479,9 @@ class WanVideoModelLoader:
patcher, _ = load_lora_for_models(patcher, None, lora_sd, lora_strength, 0)
del lora_sd
#from .utils import load_lora
#patcher = load_lora(patcher, device)
patcher.load(device, full_load=True)
patcher = apply_lora(patcher, device)
#patcher.load(device, full_load=True)
patcher.is_patched = True
del sd
@@ -1261,7 +1263,7 @@ class WanVideoSampler:
if not patcher.is_patched:
print("Patching model for control")
patcher.patch_model(device)
patcher = apply_lora(patcher, device)
patcher.is_patched = True
latent_video_length = noise.shape[1]
+33 -1
View File
@@ -29,4 +29,36 @@ def get_module_memory_mb(module):
for param in module.parameters():
if param.data is not None:
memory += param.nelement() * param.element_size()
return memory / (1024 * 1024) # Convert to MB
return memory / (1024 * 1024) # Convert to MB
def apply_lora(model, device_to=None):
to_load = []
for n, m in model.model.named_modules():
params = []
skip = False
for name, param in m.named_parameters(recurse=False):
params.append(name)
for name, param in m.named_parameters(recurse=True):
if name not in params:
skip = True # skip random weights in non leaf modules
break
if not skip and (hasattr(m, "comfy_cast_weights") or len(params) > 0):
to_load.append((n, m, params))
to_load.sort(reverse=True)
for x in to_load:
n = x[0]
m = x[1]
params = x[2]
if hasattr(m, "comfy_patched_weights"):
if m.comfy_patched_weights == True:
continue
for param in params:
model.patch_weight_to_device("{}.{}".format(n, param), device_to=device_to)
m.comfy_patched_weights = True
model.current_weight_patches_uuid = model.patches_uuid
return model