simplify LoRA loading further
This commit is contained in:
@@ -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]
|
||||
|
||||
@@ -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
|
||||
|
||||
|
||||
Reference in New Issue
Block a user