Allow merging LoRA to fp8 scaled models
This commit is contained in:
@@ -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
@@ -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:
|
||||
|
||||
@@ -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:
|
||||
|
||||
Reference in New Issue
Block a user