Fix GGUF with LoRAs and support GGUF witgh SetLoras -node
This commit is contained in:
+17
-12
@@ -33,15 +33,6 @@ def _replace_with_gguf_linear(model, compute_dtype, state_dict, prefix="", modul
|
||||
):
|
||||
key = "diffusion_model." + module_prefix + "weight"
|
||||
patch = patches.get(key, [])
|
||||
|
||||
lora_diffs = lora_strengths = None
|
||||
if len(patch) != 0:
|
||||
lora_diffs = [p[1].weights for p in patch]
|
||||
lora_strengths = [p[0] for p in patch]
|
||||
|
||||
#print("lora_diff", lora_diff)
|
||||
|
||||
#print(state_dict[module_prefix + "weight"].shape)
|
||||
in_features = state_dict[module_prefix + "weight"].shape[1]
|
||||
out_features = state_dict[module_prefix + "weight"].shape[0]
|
||||
|
||||
@@ -51,16 +42,30 @@ def _replace_with_gguf_linear(model, compute_dtype, state_dict, prefix="", modul
|
||||
in_features,
|
||||
out_features,
|
||||
module.bias is not None,
|
||||
compute_dtype=compute_dtype,
|
||||
lora_diffs=lora_diffs,
|
||||
lora_strengths = lora_strengths
|
||||
compute_dtype=compute_dtype
|
||||
)
|
||||
set_lora_params(model._modules[name], patches, module_prefix)
|
||||
model._modules[name].source_cls = type(module)
|
||||
# Force requires_grad to False to avoid unexpected errors
|
||||
model._modules[name].requires_grad_(False)
|
||||
|
||||
return model
|
||||
|
||||
def set_lora_params(module, patches, module_prefix=""):
|
||||
# Recursively set lora_diffs and lora_strengths for all GGUFLinear layers
|
||||
for name, child in module.named_children():
|
||||
child_prefix = module_prefix + name + "."
|
||||
set_lora_params(child, patches, child_prefix)
|
||||
if isinstance(module, GGUFLinear):
|
||||
key = "diffusion_model." + module_prefix + "weight"
|
||||
patch = patches.get(key, [])
|
||||
lora_diffs = lora_strengths = None
|
||||
if len(patch) != 0:
|
||||
lora_diffs = [p[1].weights for p in patch]
|
||||
lora_strengths = [p[0] for p in patch]
|
||||
module.lora_diffs = lora_diffs
|
||||
module.lora_strengths = lora_strengths
|
||||
|
||||
class GGUFLinear(nn.Linear):
|
||||
def __init__(
|
||||
self,
|
||||
|
||||
@@ -10,7 +10,7 @@ from diffusers.schedulers import FlowMatchEulerDiscreteScheduler
|
||||
from .wanvideo.modules.model import rope_params
|
||||
from .fp8_optimization import convert_linear_with_lora_and_scale, remove_lora_from_module
|
||||
from .wanvideo.schedulers import get_scheduler, get_sampling_sigmas, retrieve_timesteps, scheduler_list
|
||||
|
||||
from .gguf.gguf import set_lora_params
|
||||
from .multitalk.multitalk import timestep_transform, add_noise
|
||||
from .utils import log, print_memory, apply_lora, clip_encode_image_tiled, fourier_filter, is_image_black, add_noise_to_reference_video, optimized_scale, find_closest_valid_dim
|
||||
from .cache_methods.cache_methods import cache_report
|
||||
@@ -1278,12 +1278,16 @@ class WanVideoSampler:
|
||||
model = model.model
|
||||
transformer = model.diffusion_model
|
||||
dtype = model["dtype"]
|
||||
gguf = model["gguf"]
|
||||
control_lora = model["control_lora"]
|
||||
transformer_options = patcher.model_options.get("transformer_options", None)
|
||||
|
||||
if len(patcher.patches) != 0 and transformer_options.get("linear_with_lora", False) is True:
|
||||
log.info(f"Using {len(patcher.patches)} LoRA weight patches for WanVideo model")
|
||||
convert_linear_with_lora_and_scale(transformer, patches=patcher.patches)
|
||||
if not gguf:
|
||||
convert_linear_with_lora_and_scale(transformer, patches=patcher.patches)
|
||||
else:
|
||||
set_lora_params(transformer, patcher.patches)
|
||||
else:
|
||||
log.info("Unloading all LoRAs")
|
||||
remove_lora_from_module(transformer)
|
||||
|
||||
@@ -755,7 +755,8 @@ class WanVideoModelLoader:
|
||||
raise ValueError("Quantization should be disabled when loading GGUF models.")
|
||||
quantization = "gguf"
|
||||
gguf = True
|
||||
merge_loras = False
|
||||
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
|
||||
@@ -1223,6 +1224,7 @@ class WanVideoModelLoader:
|
||||
patcher.model["auto_cpu_offload"] = True if vram_management_args is not None else False
|
||||
patcher.model["control_lora"] = control_lora
|
||||
patcher.model["compile_args"] = compile_args
|
||||
patcher.model["gguf"] = gguf
|
||||
|
||||
if 'transformer_options' not in patcher.model_options:
|
||||
patcher.model_options['transformer_options'] = {}
|
||||
|
||||
Reference in New Issue
Block a user