Allow merging LoRA to fp8 scaled models

This commit is contained in:
kijai
2025-07-24 11:51:36 +03:00
parent d6425cab02
commit 35e637cbfd
3 changed files with 61 additions and 19 deletions
+2 -1
View File
@@ -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
+47 -11
View File
@@ -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:
+12 -7
View File
@@ -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: