Add WanVideoSetLoRAs

Node to set the LoRA weights to use with the unmerged LoRA mode, not able to merge LoRAs but allows instant LoRA switching without any loading times. The effect of unmerged LoRAs is stronger and differs from merged LoRAs.
This commit is contained in:
kijai
2025-07-21 22:39:17 +03:00
parent 41d8bd9ec9
commit 296baa30ce
4 changed files with 251 additions and 110 deletions
+44 -73
View File
@@ -27,44 +27,37 @@ def fp8_linear_forward(cls, original_dtype, input):
return cls.original_forward(input.to(original_dtype))
else:
return cls.original_forward(input)
def fp8_scaled_linear_forward(cls, original_dtype, input):
weight = cls.weight.to(original_dtype)
scale_weight = cls.scale_weight.to(input.device)
bias = cls.bias.to(original_dtype) if cls.bias is not None else None
if weight.numel() < input.numel():
weight = weight * scale_weight
else:
input = input * scale_weight
@torch.compiler.disable()
def apply_lora(weight, lora):
for lora_diff, lora_strength in zip(lora[0], lora[1]):
patch_diff = torch.mm(
lora_diff[0].flatten(start_dim=1).to(weight.device),
lora_diff[1].flatten(start_dim=1).to(weight.device)
).reshape(weight.shape)
alpha = lora_diff[2] / lora_diff[1].shape[0] if lora_diff[2] is not None else 1.0
scale = lora_strength * alpha
weight = weight.add(patch_diff, alpha=scale)
return weight
def linear_with_lora_and_scale_forward(cls, input):
# Handles both scaled and unscaled, with or without LoRA
has_scale = hasattr(cls, "scale_weight")
weight = cls.weight.to(input.dtype)
bias = cls.bias.to(input.dtype) if cls.bias is not None else None
if has_scale:
scale_weight = cls.scale_weight.to(input.device)
if weight.numel() < input.numel():
weight = weight * scale_weight
else:
input = input * scale_weight
lora = getattr(cls, "lora", None)
if lora is not None:
for lora_diff, lora_strength in zip(lora[0], lora[1]):
patch_diff = torch.mm(
lora_diff[0].flatten(start_dim=1).to(weight.device),
lora_diff[1].flatten(start_dim=1).to(weight.device)
).reshape(weight.shape)
alpha = lora_diff[2] / lora_diff[1].shape[0] if lora_diff[2] is not None else 1.0
scale = lora_strength * alpha
weight = weight.add(patch_diff, alpha=scale).to(original_dtype)
return torch.nn.functional.linear(input, weight, bias)
def linear_with_lora_forward(cls, original_dtype, input):
weight = cls.weight.to(original_dtype)
bias = cls.bias.to(original_dtype) if cls.bias is not None else None
lora = getattr(cls, "lora", None)
if lora is not None:
for lora_diff, lora_strength in zip(lora[0], lora[1]):
patch_diff = torch.mm(
lora_diff[0].flatten(start_dim=1).to(weight.device),
lora_diff[1].flatten(start_dim=1).to(weight.device)
).reshape(weight.shape)
alpha = lora_diff[2] / lora_diff[1].shape[0] if lora_diff[2] is not None else 1.0
scale = lora_strength * alpha
weight = weight.add(patch_diff, alpha=scale).to(original_dtype)
weight = apply_lora(weight, lora).to(input.dtype)
return torch.nn.functional.linear(input, weight, bias)
@@ -79,22 +72,20 @@ def convert_fp8_linear(module, original_dtype, params_to_keep={}):
setattr(submodule, "original_forward", original_forward)
setattr(submodule, "forward", lambda input, m=submodule: fp8_linear_forward(m, original_dtype, input))
def convert_fp8_scaled_linear(module, sd, original_dtype, params_to_keep={}, patches=None):
setattr(module, "fp8_scaled_enabled", True)
def convert_linear_with_lora_and_scale(module, scale_weight_keys=None, patches=None, params_to_keep={}):
for name, submodule in module.named_modules():
if not any(keyword in name for keyword in params_to_keep):
scale_key = f"{name}.scale_weight"
has_scale = scale_key in sd
weight = getattr(submodule, 'weight', None)
has_fp8_weight = weight is not None and weight.dtype in [torch.float8_e4m3fn, torch.float8_e5m2]
if has_scale:
setattr(submodule, "scale_weight", sd[scale_key])
# Set scale_weight if present
if scale_weight_keys is not None:
scale_key = f"{name}.scale_weight"
if scale_key in scale_weight_keys:
setattr(submodule, "scale_weight", scale_weight_keys[scale_key])
# Set LoRA if present
if patches is not None:
patch_key = f"diffusion_model.{name}.weight"
patch = patches.get(patch_key, [])
#print("Patches for", patch_key, ":", patch)
if len(patch) != 0:
lora_diffs = []
for p in patch:
@@ -105,35 +96,15 @@ def convert_fp8_scaled_linear(module, sd, original_dtype, params_to_keep={}, pat
lora_diffs.append(lora_obj[1])
else:
continue
lora_strengths = [p[0] for p in patch]
lora = (lora_diffs, lora_strengths)
setattr(submodule, "lora", lora)
if isinstance(submodule, nn.Linear) and (has_scale and has_fp8_weight):
original_forward = submodule.forward
setattr(submodule, "original_forward", original_forward)
setattr(submodule, "forward", lambda input, m=submodule: fp8_scaled_linear_forward(m, original_dtype, input))
def convert_linear_with_lora(module, original_dtype, patches=None):
for name, submodule in module.named_modules():
if isinstance(submodule, nn.Linear):
patch_key = f"diffusion_model.{name}.weight"
patch = patches.get(patch_key, [])
if len(patch) != 0:
lora_diffs = []
for p in patch:
lora_obj = p[1]
if hasattr(lora_obj, "weights"):
lora_diffs.append(lora_obj.weights)
elif isinstance(lora_obj, tuple) and lora_obj[0] == "diff":
lora_diffs.append(lora_obj[1])
else:
continue
lora_strengths = [p[0] for p in patch]
lora = (lora_diffs, lora_strengths)
setattr(submodule, "lora", lora)
# original_forward = submodule.forward
# setattr(submodule, "original_forward", original_forward)
setattr(submodule, "forward", lambda input, m=submodule: linear_with_lora_forward(m, original_dtype, input))
# Set forward if Linear and has either scale or lora
if isinstance(submodule, nn.Linear):
has_scale = hasattr(submodule, "scale_weight")
has_lora = hasattr(submodule, "lora")
if has_scale or has_lora:
original_forward = submodule.forward
setattr(submodule, "original_forward", original_forward)
setattr(submodule, "forward", lambda input, m=submodule: linear_with_lora_and_scale_forward(m, input))
+28 -3
View File
@@ -8,7 +8,7 @@ import inspect
from diffusers.schedulers import FlowMatchEulerDiscreteScheduler
from .wanvideo.modules.model import rope_params
from .fp8_optimization import convert_linear_with_lora_and_scale
from .wanvideo.schedulers import get_scheduler, get_sampling_sigmas, retrieve_timesteps, scheduler_list
from .multitalk.multitalk import timestep_transform, add_noise
@@ -1174,12 +1174,37 @@ class WanVideoSampler:
transformer = model.diffusion_model
dtype = model["dtype"]
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)} patches for WanVideo model")
convert_linear_with_lora_and_scale(transformer, patches=patcher.patches)
#compile
compile_args = model["compile_args"]
if compile_args is not None and model["auto_cpu_offload"] is False:
torch._dynamo.config.cache_size_limit = compile_args["dynamo_cache_size_limit"]
try:
if hasattr(torch, '_dynamo') and hasattr(torch._dynamo, 'config'):
torch._dynamo.config.recompile_limit = compile_args["dynamo_recompile_limit"]
except Exception as e:
log.warning(f"Could not set recompile_limit: {e}")
if compile_args["compile_transformer_blocks_only"]:
for i, block in enumerate(transformer.blocks):
if hasattr(block, "_orig_mod"):
block = block._orig_mod
transformer.blocks[i] = torch.compile(block, fullgraph=compile_args["fullgraph"], dynamic=compile_args["dynamic"], backend=compile_args["backend"], mode=compile_args["mode"])
if transformer.vace_layers is not None:
for i, block in enumerate(transformer.vace_blocks):
if hasattr(block, "_orig_mod"):
block = block._orig_mod
transformer.vace_blocks[i] = torch.compile(block, fullgraph=compile_args["fullgraph"], dynamic=compile_args["dynamic"], backend=compile_args["backend"], mode=compile_args["mode"])
else:
transformer = torch.compile(transformer, fullgraph=compile_args["fullgraph"], dynamic=compile_args["dynamic"], backend=compile_args["backend"], mode=compile_args["mode"])
multitalk_sampling = image_embeds.get("multitalk_sampling", False)
if not multitalk_sampling and scheduler == "multitalk":
raise Exception("multitalk scheduler is only for multitalk sampling when using ImagetoVideoMultiTalk -node")
transformer_options = patcher.model_options.get("transformer_options", None)
steps = int(steps/denoise_strength)
+177 -33
View File
@@ -1,5 +1,5 @@
import torch
import os, gc
import os, gc, uuid
from .utils import log, apply_lora
import numpy as np
from tqdm import tqdm
@@ -11,6 +11,8 @@ from .wanvideo.modules.clip import CLIPModel
from accelerate import init_empty_weights
from accelerate.utils import set_module_tensor_to_device
from .fp8_optimization import convert_linear_with_lora_and_scale
import folder_paths
import comfy.model_management as mm
from comfy.utils import load_torch_file, ProgressBar
@@ -543,6 +545,153 @@ class WanVideoLoraBlockEdit:
}
return (selected,)
def model_lora_keys_unet(model, key_map={}):
sd = model.state_dict()
sdk = sd.keys()
for k in sdk:
k = k.replace("_orig_mod.", "")
if k.startswith("diffusion_model."):
if k.endswith(".weight"):
key_lora = k[len("diffusion_model."):-len(".weight")].replace(".", "_")
key_map["lora_unet_{}".format(key_lora)] = k
key_map["{}".format(k[:-len(".weight")])] = k #generic lora format without any weird key names
else:
key_map["{}".format(k)] = k #generic lora format for not .weight without any weird key names
diffusers_keys = comfy.utils.unet_to_diffusers(model.model_config.unet_config)
for k in diffusers_keys:
if k.endswith(".weight"):
unet_key = "diffusion_model.{}".format(diffusers_keys[k])
key_lora = k[:-len(".weight")].replace(".", "_")
key_map["lora_unet_{}".format(key_lora)] = unet_key
key_map["lycoris_{}".format(key_lora)] = unet_key #simpletuner lycoris format
diffusers_lora_prefix = ["", "unet."]
for p in diffusers_lora_prefix:
diffusers_lora_key = "{}{}".format(p, k[:-len(".weight")].replace(".to_", ".processor.to_"))
if diffusers_lora_key.endswith(".to_out.0"):
diffusers_lora_key = diffusers_lora_key[:-2]
key_map[diffusers_lora_key] = unet_key
return key_map
def add_patches(patcher, patches, strength_patch=1.0, strength_model=1.0):
with patcher.use_ejected():
p = set()
model_sd = patcher.model.state_dict()
for k in patches:
offset = None
function = None
if isinstance(k, str):
key = k
else:
offset = k[1]
key = k[0]
if len(k) > 2:
function = k[2]
# Check for key, or key with '._orig_mod' inserted after block number, in model_sd
key_in_sd = key in model_sd
key_orig_mod = None
if not key_in_sd:
# Try to insert '._orig_mod' after the block number if pattern matches
parts = key.split('.')
# Look for 'blocks', block number, then insert
try:
idx = parts.index('blocks')
if idx + 1 < len(parts):
# Only if the next part is a number
if parts[idx+1].isdigit():
new_parts = parts[:idx+2] + ['_orig_mod'] + parts[idx+2:]
key_orig_mod = '.'.join(new_parts)
except ValueError:
pass
key_orig_mod_in_sd = key_orig_mod is not None and key_orig_mod in model_sd
if key_in_sd or key_orig_mod_in_sd:
actual_key = key if key_in_sd else key_orig_mod
p.add(k)
current_patches = patcher.patches.get(actual_key, [])
current_patches.append((strength_patch, patches[k], strength_model, offset, function))
patcher.patches[actual_key] = current_patches
patcher.patches_uuid = uuid.uuid4()
return list(p)
def load_lora_for_models_mod(model, lora, strength_model):
key_map = {}
if model is not None:
key_map = model_lora_keys_unet(model.model, key_map)
loaded = comfy.lora.load_lora(lora, key_map)
new_modelpatcher = model.clone()
k = add_patches(new_modelpatcher, loaded, strength_model)
k = set(k)
for x in loaded:
if (x not in k):
log.warning("NOT LOADED {}".format(x))
return (new_modelpatcher)
class WanVideoSetLoRAs:
@classmethod
def INPUT_TYPES(s):
return {
"required":
{
"model": ("WANVIDEOMODEL", ),
"lora": ("WANVIDLORA", ),
},
}
RETURN_TYPES = ("WANVIDEOMODEL",)
RETURN_NAMES = ("model", )
FUNCTION = "setlora"
CATEGORY = "WanVideoWrapper"
EXPERIMENTAL = True
DESCRIPTION = "Sets the LoRA weights to be used directly in linear layers of the model, this does NOT merge LoRAs"
def setlora(self, model, lora):
if lora is None:
return (model,)
patcher = model.clone()
lora_low_mem_load = merge_loras = False
for l in lora:
lora_low_mem_load = l.get("low_mem_load", False)
merge_loras = l.get("merge_loras", True)
if lora_low_mem_load is True or merge_loras is True:
raise ValueError("Set LoRA node does not use low_mem_load and can't merege LoRAs")
for l in lora:
log.info(f"Loading LoRA: {l['name']} with strength: {l['strength']}")
lora_path = l["path"]
lora_strength = l["strength"]
lora_sd = load_torch_file(lora_path, safe_load=True)
if "dwpose_embedding.0.weight" in lora_sd: #unianimate
raise NotImplementedError("Unianimate LoRA patching is not implemented in this node.")
lora_sd = standardize_lora_key_format(lora_sd)
if l["blocks"]:
lora_sd = filter_state_dict_by_blocks(lora_sd, l["blocks"], l.get("layer_filter", []))
if "diffusion_model.patch_embedding.lora_A.weight" in lora_sd:
raise NotImplementedError("Control LoRA patching is not implemented in this node.")
patcher = load_lora_for_models_mod(patcher, lora_sd, lora_strength)
del lora_sd
if 'transformer_options' not in patcher.model_options:
patcher.model_options['transformer_options'] = {}
patcher.model_options['transformer_options']["linear_with_lora"] = True
return (patcher,)
#region Model loading
class WanVideoModelLoader:
@classmethod
@@ -600,9 +749,6 @@ class WanVideoModelLoader:
except Exception as e:
raise ValueError(f"Can't import SageAttention: {str(e)}")
device = mm.get_torch_device()
offload_device = mm.unet_offload_device()
gguf = False
if model.endswith(".gguf"):
if quantization != "disabled":
@@ -972,8 +1118,7 @@ class WanVideoModelLoader:
patcher.model.is_patched = True
if "fast" in quantization:
if not merge_loras:
raise ValueError("FP8 fast quantization requires LoRAs to be merged into the model, please set merge_loras=True in the LoRA input")
@@ -984,13 +1129,15 @@ class WanVideoModelLoader:
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
log.info("Using FP8 scaled linear quantization")
from .fp8_optimization import convert_fp8_scaled_linear
convert_fp8_scaled_linear(patcher.model.diffusion_model, sd, base_dtype, params_to_keep=params_to_keep, patches=patcher.patches)
elif not merge_loras and not gguf:
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:
log.info("LoRAs will be applied at runtime")
from .fp8_optimization import convert_linear_with_lora
convert_linear_with_lora(patcher.model.diffusion_model, base_dtype, patches=patcher.patches)
convert_linear_with_lora_and_scale(patcher.model.diffusion_model, patches=patcher.patches)
del sd
@@ -1040,22 +1187,22 @@ class WanVideoModelLoader:
compile_args = compile_args,
)
#compile
if compile_args is not None and vram_management_args is None:
torch._dynamo.config.cache_size_limit = compile_args["dynamo_cache_size_limit"]
try:
if hasattr(torch, '_dynamo') and hasattr(torch._dynamo, 'config'):
torch._dynamo.config.recompile_limit = compile_args["dynamo_recompile_limit"]
except Exception as e:
log.warning(f"Could not set recompile_limit: {e}")
if compile_args["compile_transformer_blocks_only"]:
for i, block in enumerate(patcher.model.diffusion_model.blocks):
patcher.model.diffusion_model.blocks[i] = torch.compile(block, fullgraph=compile_args["fullgraph"], dynamic=compile_args["dynamic"], backend=compile_args["backend"], mode=compile_args["mode"])
if vace_layers is not None:
for i, block in enumerate(patcher.model.diffusion_model.vace_blocks):
patcher.model.diffusion_model.vace_blocks[i] = torch.compile(block, fullgraph=compile_args["fullgraph"], dynamic=compile_args["dynamic"], backend=compile_args["backend"], mode=compile_args["mode"])
else:
patcher.model.diffusion_model = torch.compile(patcher.model.diffusion_model, fullgraph=compile_args["fullgraph"], dynamic=compile_args["dynamic"], backend=compile_args["backend"], mode=compile_args["mode"])
# #compile
# if compile_args is not None and vram_management_args is None:
# torch._dynamo.config.cache_size_limit = compile_args["dynamo_cache_size_limit"]
# try:
# if hasattr(torch, '_dynamo') and hasattr(torch._dynamo, 'config'):
# torch._dynamo.config.recompile_limit = compile_args["dynamo_recompile_limit"]
# except Exception as e:
# log.warning(f"Could not set recompile_limit: {e}")
# if compile_args["compile_transformer_blocks_only"]:
# for i, block in enumerate(patcher.model.diffusion_model.blocks):
# patcher.model.diffusion_model.blocks[i] = torch.compile(block, fullgraph=compile_args["fullgraph"], dynamic=compile_args["dynamic"], backend=compile_args["backend"], mode=compile_args["mode"])
# if vace_layers is not None:
# for i, block in enumerate(patcher.model.diffusion_model.vace_blocks):
# patcher.model.diffusion_model.vace_blocks[i] = torch.compile(block, fullgraph=compile_args["fullgraph"], dynamic=compile_args["dynamic"], backend=compile_args["backend"], mode=compile_args["mode"])
# else:
# patcher.model.diffusion_model = torch.compile(patcher.model.diffusion_model, fullgraph=compile_args["fullgraph"], dynamic=compile_args["dynamic"], backend=compile_args["backend"], mode=compile_args["mode"])
if load_device == "offload_device" and patcher.model.diffusion_model.device != offload_device:
log.info(f"Moving diffusion model from {patcher.model.diffusion_model.device} to {offload_device}")
@@ -1070,6 +1217,7 @@ class WanVideoModelLoader:
patcher.model["quantization"] = quantization
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
if 'transformer_options' not in patcher.model_options:
patcher.model_options['transformer_options'] = {}
@@ -1106,9 +1254,6 @@ class WanVideoVAELoader:
def loadmodel(self, model_name, precision):
from .wanvideo.wan_video_vae import WanVideoVAE
device = mm.get_torch_device()
offload_device = mm.unet_offload_device()
dtype = {"bf16": torch.bfloat16, "fp16": torch.float16, "fp32": torch.float32}[precision]
#with open(os.path.join(script_directory, 'configs', 'hy_vae_config.json')) as f:
# vae_config = json.load(f)
@@ -1149,9 +1294,6 @@ class WanVideoTinyVAELoader:
def loadmodel(self, model_name, precision, parallel=False):
from .taehv import TAEHV
device = mm.get_torch_device()
offload_device = mm.unet_offload_device()
dtype = {"bf16": torch.bfloat16, "fp16": torch.float16, "fp32": torch.float32}[precision]
model_path = folder_paths.get_full_path("vae_approx", model_name)
vae_sd = load_torch_file(model_path, safe_load=True)
@@ -1311,6 +1453,7 @@ NODE_CLASS_MAPPINGS = {
"WanVideoModelLoader": WanVideoModelLoader,
"WanVideoVAELoader": WanVideoVAELoader,
"WanVideoLoraSelect": WanVideoLoraSelect,
"WanVideoSetLoRAs": WanVideoSetLoRAs,
"WanVideoLoraBlockEdit": WanVideoLoraBlockEdit,
"WanVideoTinyVAELoader": WanVideoTinyVAELoader,
"WanVideoVACEModelSelect": WanVideoVACEModelSelect,
@@ -1326,6 +1469,7 @@ NODE_DISPLAY_NAME_MAPPINGS = {
"WanVideoModelLoader": "WanVideo Model Loader",
"WanVideoVAELoader": "WanVideo VAE Loader",
"WanVideoLoraSelect": "WanVideo Lora Select",
"WanVideoSetLoRAs": "WanVideo Set LoRAs",
"WanVideoLoraBlockEdit": "WanVideo Lora Block Edit",
"WanVideoTinyVAELoader": "WanVideo Tiny VAE Loader",
"WanVideoVACEModelSelect": "WanVideo VACE Module Select",
+2 -1
View File
@@ -108,7 +108,8 @@ def apply_lora(model, device_to, transformer_load_device, params_to_keep=None, d
else:
model.patch_weight_to_device("{}.{}".format(name, param), device_to=device_to)
model.backup["{}.{}".format(name, param)] = None
set_module_tensor_to_device(m, param, device=transformer_load_device)
if device_to != transformer_load_device:
set_module_tensor_to_device(m, param, device=transformer_load_device)
if low_mem_load:
try:
set_module_tensor_to_device(model.model.diffusion_model, key, device=transformer_load_device, dtype=dtype_to_use, value=model.model.diffusion_model.state_dict()[key])