diff --git a/gguf/gguf.py b/gguf/gguf.py index aa6288c..eed3435 100644 --- a/gguf/gguf.py +++ b/gguf/gguf.py @@ -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, diff --git a/nodes.py b/nodes.py index 9d733fc..a83803b 100644 --- a/nodes.py +++ b/nodes.py @@ -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) diff --git a/nodes_model_loading.py b/nodes_model_loading.py index 0fb1b12..0b0a81f 100644 --- a/nodes_model_loading.py +++ b/nodes_model_loading.py @@ -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'] = {}