From 35e637cbfd3f3a119ab629d68338772d8385693e Mon Sep 17 00:00:00 2001 From: kijai <40791699+kijai@users.noreply.github.com> Date: Thu, 24 Jul 2025 11:51:36 +0300 Subject: [PATCH] Allow merging LoRA to fp8 scaled models --- nodes.py | 3 ++- nodes_model_loading.py | 58 ++++++++++++++++++++++++++++++++++-------- utils.py | 19 +++++++++----- 3 files changed, 61 insertions(+), 19 deletions(-) diff --git a/nodes.py b/nodes.py index bae409c..9d733fc 100644 --- a/nodes.py +++ b/nodes.py @@ -1282,9 +1282,10 @@ class WanVideoSampler: 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)} patches for WanVideo model") + log.info(f"Using {len(patcher.patches)} LoRA weight patches for WanVideo model") convert_linear_with_lora_and_scale(transformer, patches=patcher.patches) else: + log.info("Unloading all LoRAs") remove_lora_from_module(transformer) #compile diff --git a/nodes_model_loading.py b/nodes_model_loading.py index 77e68cb..0fb1b12 100644 --- a/nodes_model_loading.py +++ b/nodes_model_loading.py @@ -801,9 +801,6 @@ class WanVideoModelLoader: if "scaled_fp8" in sd and "scaled" not in quantization: raise ValueError("The model is a scaled fp8 model, please set quantization to '_scaled'") - if merge_loras and "scaled" in quantization: - raise ValueError("scaled models currently do not support merging LoRAs, please disable merging or use a non-scaled model") - if "vace_blocks.0.after_proj.weight" in sd and not "patch_embedding.weight" in sd: raise ValueError("You are attempting to load a VACE module as a WanVideo model, instead you should use the vace_model input and matching T2V base model") @@ -1033,6 +1030,13 @@ class WanVideoModelLoader: patcher.model.is_patched = False control_lora = False + + scale_weights = {} + if "scaled" in quantization: + scale_weights = {} + for k, v in sd.items(): + if k.endswith(".scale_weight"): + scale_weights[k] = v if lora is not None: for l in lora: @@ -1087,9 +1091,12 @@ class WanVideoModelLoader: del lora_sd - if not gguf and not "scaled" in quantization and merge_loras: + if not gguf 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, control_lora=control_lora) + 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, scale_weights=scale_weights) if gguf: #from diffusers.quantizers.gguf.utils import _replace_with_gguf_linear, GGUFParameter @@ -1130,11 +1137,7 @@ class WanVideoModelLoader: print(params_to_keep) convert_fp8_linear(patcher.model.diffusion_model, base_dtype, params_to_keep=params_to_keep) - if "scaled" in quantization: - scale_weights = {} - for k, v in sd.items(): - if k.endswith(".scale_weight"): - scale_weights[k] = v + if "scaled" in quantization and not merge_loras: log.info("Using FP8 scaled linear quantization") convert_linear_with_lora_and_scale(patcher.model.diffusion_model, scale_weights, params_to_keep=params_to_keep, patches=patcher.patches) elif lora is not None and not merge_loras and not gguf: @@ -1223,7 +1226,8 @@ class WanVideoModelLoader: if 'transformer_options' not in patcher.model_options: patcher.model_options['transformer_options'] = {} - patcher.model_options["transformer_options"]["block_swap_args"] = block_swap_args + patcher.model_options["transformer_options"]["block_swap_args"] = block_swap_args + patcher.model_options["transformer_options"]["linear_with_lora"] = True if not merge_loras else False for model in mm.current_loaded_models: if model._model() == patcher: @@ -1231,6 +1235,38 @@ class WanVideoModelLoader: return (patcher,) +# class WanVideoSaveModel: +# @classmethod +# def INPUT_TYPES(s): +# return { +# "required": { +# "model": ("WANVIDEOMODEL", {"tooltip": "WANVideo model to save"}), +# "output_path": ("STRING", {"default": "", "multiline": False, "tooltip": "Path to save the model"}), +# }, +# } + +# RETURN_TYPES = () +# FUNCTION = "savemodel" +# CATEGORY = "WanVideoWrapper" +# DESCRIPTION = "Saves the model including merged LoRAs and quantization to diffusion_models/WanVideoWrapperSavedModels" +# OUTPUT_NODE = True + +# def savemodel(self, model, output_path): +# from safetensors.torch import save_file +# model_sd = model.model.diffusion_model.state_dict() +# for k in model_sd.keys(): +# print("key:", k, "shape:", model_sd[k].shape, "dtype:", model_sd[k].dtype, "device:", model_sd[k].device) +# model_sd +# model_name = os.path.basename(model.model["model_name"]) +# if not output_path: +# output_path = os.path.join(folder_paths.models_dir, "diffusion_models", "WanVideoWrapperSavedModels", "saved_" + model_name) +# else: +# output_path = os.path.join(output_path, model_name) +# log.info(f"Saving model to {output_path}") +# os.makedirs(os.path.dirname(output_path), exist_ok=True) +# save_file(model_sd, output_path) +# return () + #region load VAE class WanVideoVAELoader: diff --git a/utils.py b/utils.py index 9a1c13e..6fbc5c0 100644 --- a/utils.py +++ b/utils.py @@ -42,7 +42,7 @@ 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, backup_keys=False): +def patch_weight_to_device(self, key, device_to=None, inplace_update=False, backup_keys=False, scale_weight=None): if key not in self.patches: return @@ -59,7 +59,11 @@ def patch_weight_to_device(self, key, device_to=None, inplace_update=False, back if convert_func is not None: temp_weight = convert_func(temp_weight, inplace=True) + if scale_weight is not None: + temp_weight = temp_weight * scale_weight.to(temp_weight.device, temp_weight.dtype) + out_weight = calculate_weight(self.patches[key], temp_weight, key) + if set_func is None: out_weight = stochastic_rounding(out_weight, weight.dtype, seed=string_to_seed(key)) if inplace_update: @@ -69,7 +73,7 @@ def patch_weight_to_device(self, key, device_to=None, inplace_update=False, back 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, control_lora=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, scale_weights={}): model.patch_weight_to_device = types.MethodType(patch_weight_to_device, model) to_load = [] for n, m in model.model.named_modules(): @@ -99,17 +103,18 @@ def apply_lora(model, device_to, transformer_load_device, params_to_keep=None, d dtype_to_use = base_dtype if any(keyword in name for keyword in params_to_keep) else dtype if "patch_embedding" in name: dtype_to_use = torch.float32 - if name.startswith("diffusion_model."): - name_no_prefix = name[len("diffusion_model."):] - key = "{}.{}".format(name_no_prefix, param) + key = f"{name.replace('diffusion_model.', '')}.{param}" try: set_module_tensor_to_device(model.model.diffusion_model, key, device=transformer_load_device, dtype=dtype_to_use, value=state_dict[key]) except: continue + key = f"{name}.{param}" + if scale_weights is not None: + scale_key = key.replace("weight", "scale_weight").replace("diffusion_model.", "") if "weight" in key else None if low_mem_load: - model.patch_weight_to_device("{}.{}".format(name, param), device_to=device_to, inplace_update=True, backup_keys=control_lora) + model.patch_weight_to_device(f"{name}.{param}", device_to=device_to, inplace_update=True, backup_keys=control_lora, scale_weight=scale_weights.get(scale_key, None)) else: - model.patch_weight_to_device("{}.{}".format(name, param), device_to=device_to, backup_keys=control_lora) + model.patch_weight_to_device(f"{name}.{param}", device_to=device_to, backup_keys=control_lora, scale_weight=scale_weights.get(scale_key, None)) if device_to != transformer_load_device: set_module_tensor_to_device(m, param, device=transformer_load_device) if low_mem_load: