Make lora torch.compile optional for unmerged lora application

This change has caused issues especially with LoRAs that have dynamic rank. Will now be disabled by default, to allow full graph with unmerged LoRAs the option to allow compile is available in the Torch Compile Settings -node
This commit is contained in:
kijai
2025-11-04 01:34:51 +02:00
parent 8ce6916d72
commit 9fa4140159
5 changed files with 78 additions and 43 deletions
+41 -24
View File
@@ -3,15 +3,20 @@ import torch.nn as nn
from accelerate import init_empty_weights
#based on https://github.com/huggingface/diffusers/blob/main/src/diffusers/quantizers/gguf/utils.py
def _replace_linear(model, compute_dtype, state_dict, prefix="", patches=None, scale_weights=None):
def _replace_linear(model, compute_dtype, state_dict, prefix="", patches=None, scale_weights=None, compile_args=None):
has_children = list(model.children())
if not has_children:
return
allow_compile = False
for name, module in model.named_children():
if compile_args is not None:
allow_compile = compile_args.get("allow_unmerged_lora_compile", False)
module_prefix = prefix + name + "."
module_prefix = module_prefix.replace("_orig_mod.", "")
_replace_linear(module, compute_dtype, state_dict, module_prefix, patches, scale_weights)
_replace_linear(module, compute_dtype, state_dict, module_prefix, patches, scale_weights, compile_args)
if isinstance(module, nn.Linear) and "loras" not in module_prefix:
in_features = state_dict[module_prefix + "weight"].shape[1]
@@ -25,7 +30,8 @@ def _replace_linear(model, compute_dtype, state_dict, prefix="", patches=None, s
out_features,
module.bias is not None,
compute_dtype=compute_dtype,
scale_weight=scale_weights.get(scale_key) if scale_weights else None
scale_weight=scale_weights.get(scale_key) if scale_weights else None,
allow_compile=allow_compile
)
model._modules[name].source_cls = type(module)
model._modules[name].requires_grad_(False)
@@ -77,7 +83,8 @@ class CustomLinear(nn.Linear):
bias=False,
compute_dtype=None,
device=None,
scale_weight=None
scale_weight=None,
allow_compile=False
) -> None:
super().__init__(in_features, out_features, bias, device)
self.compute_dtype = compute_dtype
@@ -85,6 +92,10 @@ class CustomLinear(nn.Linear):
self.step = 0
self.scale_weight = scale_weight
self.lora_strengths = []
self.allow_compile = allow_compile
if not allow_compile:
self._get_weight_with_lora = torch.compiler.disable()(self._get_weight_with_lora)
def set_lora_diffs(self, lora_diffs, device=torch.device("cpu")):
self.lora_diffs = []
@@ -98,25 +109,11 @@ class CustomLinear(nn.Linear):
self.register_buffer(f"lora_diff_{i}_0", diff[0].to(device))
self.lora_diffs.append(f"lora_diff_{i}_0")
def forward(self, input):
if self.bias is not None:
bias = self.bias.to(input)
else:
bias = None
weight = self.weight.to(input)
if self.scale_weight is not None:
if weight.numel() < input.numel():
weight = weight * self.scale_weight
else:
input = input * self.scale_weight
if hasattr(self, f"lora_diff_0_0"):
weight = self.apply_lora(weight).to(self.compute_dtype)
return torch.nn.functional.linear(input, weight, bias)
def apply_lora(self, weight):
def _get_weight_with_lora(self, weight):
"""Apply LoRA outside compiled region"""
if not hasattr(self, "lora_diff_0_0"):
return weight
for lora_diff_names, lora_strength in zip(self.lora_diffs, self.lora_strengths):
if isinstance(lora_strength, list):
lora_strength = lora_strength[self.step]
@@ -131,7 +128,7 @@ class CustomLinear(nn.Linear):
patch_diff = torch.mm(
lora_diff_0.flatten(start_dim=1),
lora_diff_1.flatten(start_dim=1)
).reshape(weight.shape) + 0
).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)
@@ -139,6 +136,26 @@ class CustomLinear(nn.Linear):
lora_diff = getattr(self, lora_diff_names)
weight = weight.add(lora_diff, alpha=lora_strength)
return weight
def forward(self, input):
if self.bias is not None:
bias = self.bias.to(input)
else:
bias = None
weight = self.weight.to(input)
if self.scale_weight is not None:
if weight.numel() < input.numel():
weight = weight * self.scale_weight
else:
input = input * self.scale_weight
weight = self._get_weight_with_lora(weight)
if self.compute_dtype is not None:
weight = weight.to(self.compute_dtype)
return torch.nn.functional.linear(input, weight, bias)
def remove_lora_from_module(module):
for name, submodule in module.named_modules():
+23 -10
View File
@@ -19,7 +19,7 @@ def load_gguf(model_path):
return parsed_parameters, reader
#based on https://github.com/huggingface/diffusers/blob/main/src/diffusers/quantizers/gguf/utils.py
def _replace_with_gguf_linear(model, compute_dtype, state_dict, prefix="", modules_to_not_convert=[], patches=None):
def _replace_with_gguf_linear(model, compute_dtype, state_dict, prefix="", modules_to_not_convert=[], patches=None, compile_args=None):
def _should_convert_to_gguf(state_dict, prefix):
weight_key = prefix + "weight"
return weight_key in state_dict and isinstance(state_dict[weight_key], GGUFParameter)
@@ -27,10 +27,14 @@ def _replace_with_gguf_linear(model, compute_dtype, state_dict, prefix="", modul
has_children = list(model.children())
if not has_children:
return
allow_compile = False
for name, module in model.named_children():
if compile_args is not None:
allow_compile = compile_args.get("allow_unmerged_lora_compile", False)
module_prefix = prefix + name + "."
_replace_with_gguf_linear(module, compute_dtype, state_dict, module_prefix, modules_to_not_convert, patches)
_replace_with_gguf_linear(module, compute_dtype, state_dict, module_prefix, modules_to_not_convert, patches, compile_args)
if (
isinstance(module, nn.Linear)
@@ -46,7 +50,8 @@ 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
compute_dtype=compute_dtype,
allow_compile=allow_compile
)
model._modules[name].source_cls = type(module)
@@ -95,19 +100,23 @@ class GGUFLinear(nn.Linear):
bias=False,
compute_dtype=None,
device=None,
allow_compile=False
) -> None:
super().__init__(in_features, out_features, bias, device)
self.compute_dtype = compute_dtype
self.lora_diffs = []
self.lora_strengths = []
self.step = 0
self.allow_compile = allow_compile
if not allow_compile:
self._get_weight_with_lora = torch.compiler.disable()(self._get_weight_with_lora)
def forward(self, inputs):
weight = dequantize_gguf_tensor(self.weight).to(self.compute_dtype)
bias = self.bias.to(self.compute_dtype) if self.bias is not None else None
if hasattr(self, f"lora_diff_0_0"):
weight = self.apply_lora(weight).to(self.compute_dtype)
weight = self._get_weight_with_lora(weight)#.to(self.compute_dtype)
return torch.nn.functional.linear(inputs, weight, bias)
@@ -115,15 +124,19 @@ class GGUFLinear(nn.Linear):
self.lora_diffs = []
for i, diff in enumerate(lora_diffs):
if len(diff) > 1:
self.register_buffer(f"lora_diff_{i}_0", diff[0].to(device))
self.register_buffer(f"lora_diff_{i}_1", diff[1].to(device))
self.register_buffer(f"lora_diff_{i}_0", diff[0].to(device, self.compute_dtype))
self.register_buffer(f"lora_diff_{i}_1", diff[1].to(device, self.compute_dtype))
setattr(self, f"lora_diff_{i}_2", diff[2])
self.lora_diffs.append((f"lora_diff_{i}_0", f"lora_diff_{i}_1", f"lora_diff_{i}_2"))
else:
self.register_buffer(f"lora_diff_{i}_0", diff[0].to(device))
self.register_buffer(f"lora_diff_{i}_0", diff[0].to(device, self.compute_dtype))
self.lora_diffs.append(f"lora_diff_{i}_0")
def apply_lora(self, weight):
def _get_weight_with_lora(self, weight):
"""Apply LoRA outside compiled region"""
if not hasattr(self, "lora_diff_0_0"):
return weight
for lora_diff_names, lora_strength in zip(self.lora_diffs, self.lora_strengths):
if isinstance(lora_strength, list):
lora_strength = lora_strength[self.step]
@@ -138,7 +151,7 @@ class GGUFLinear(nn.Linear):
patch_diff = torch.mm(
lora_diff_0.flatten(start_dim=1),
lora_diff_1.flatten(start_dim=1)
).reshape(weight.shape) + 0
).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)
+7 -4
View File
@@ -330,6 +330,7 @@ class WanVideoTorchCompileSettings:
"optional": {
"dynamo_recompile_limit": ("INT", {"default": 128, "min": 0, "max": 1024, "step": 1, "tooltip": "torch._dynamo.config.recompile_limit"}),
"force_parameter_static_shapes": ("BOOLEAN", {"default": False, "tooltip": "torch._dynamo.config.force_parameter_static_shapes"}),
"allow_unmerged_lora_compile": ("BOOLEAN", {"default": False, "tooltip": "Allow LoRA application to be compiled with torch.compile to avoid graph breaks, causes issues with some LoRAs, mostly dynamic ones"}),
},
}
RETURN_TYPES = ("WANCOMPILEARGS",)
@@ -338,7 +339,8 @@ class WanVideoTorchCompileSettings:
CATEGORY = "WanVideoWrapper"
DESCRIPTION = "torch.compile settings, when connected to the model loader, torch.compile of the selected layers is attempted. Requires Triton and torch > 2.7.0 is recommended"
def set_args(self, backend, fullgraph, mode, dynamic, dynamo_cache_size_limit, compile_transformer_blocks_only, dynamo_recompile_limit=128, force_parameter_static_shapes=True):
def set_args(self, backend, fullgraph, mode, dynamic, dynamo_cache_size_limit, compile_transformer_blocks_only, dynamo_recompile_limit=128,
force_parameter_static_shapes=True, allow_unmerged_lora_compile=False):
compile_args = {
"backend": backend,
@@ -349,6 +351,7 @@ class WanVideoTorchCompileSettings:
"dynamo_recompile_limit": dynamo_recompile_limit,
"compile_transformer_blocks_only": compile_transformer_blocks_only,
"force_parameter_static_shapes": force_parameter_static_shapes,
"allow_unmerged_lora_compile": allow_unmerged_lora_compile,
}
return (compile_args, )
@@ -782,7 +785,7 @@ def rename_fuser_block(name):
return new_name
def load_weights(transformer, sd=None, weight_dtype=None, base_dtype=None,
transformer_load_device=None, block_swap_args=None, gguf=False, reader=None, patcher=None):
transformer_load_device=None, block_swap_args=None, gguf=False, reader=None, patcher=None, compile_args=None):
params_to_keep = {"time_in", "patch_embedding", "time_", "modulation", "text_embedding",
"adapter", "add", "ref_conv", "casual_audio_encoder", "cond_encoder", "frame_packer", "audio_proj_glob", "face_encoder", "fuser_block"}
param_count = sum(1 for _ in transformer.named_parameters())
@@ -837,7 +840,7 @@ def load_weights(transformer, sd=None, weight_dtype=None, base_dtype=None,
if not getattr(transformer, "gguf_patched", False):
transformer = _replace_with_gguf_linear(
transformer, base_dtype, sd, patches=patcher.patches
transformer, base_dtype, sd, patches=patcher.patches, compile_args=compile_args
)
transformer.gguf_patched = True
else:
@@ -1534,7 +1537,7 @@ class WanVideoModelLoader:
transformer.patched_linear = False
sd = None
elif "scaled" in quantization or lora is not None:
transformer = _replace_linear(transformer, base_dtype, sd, scale_weights=scale_weights)
transformer = _replace_linear(transformer, base_dtype, sd, scale_weights=scale_weights, compile_args=compile_args)
transformer.patched_linear = True
if "fast" in quantization:
+6 -4
View File
@@ -194,14 +194,16 @@ class WanVideoSampler:
if hasattr(block, 'audio_block'):
block.audio_block = None
if not transformer.patched_linear and patcher.model["sd"] is not None and len(patcher.patches) != 0:
transformer = _replace_linear(transformer, dtype, patcher.model["sd"])
if not transformer.patched_linear and patcher.model["sd"] is not None and len(patcher.patches) != 0 and gguf_reader is None:
transformer = _replace_linear(transformer, dtype, patcher.model["sd"], model["compile_args"])
transformer.patched_linear = True
if patcher.model["sd"] is not None and gguf_reader is None:
load_weights(patcher.model.diffusion_model, patcher.model["sd"], weight_dtype, base_dtype=dtype, transformer_load_device=device, block_swap_args=block_swap_args)
load_weights(patcher.model.diffusion_model, patcher.model["sd"], weight_dtype, base_dtype=dtype, transformer_load_device=device,
block_swap_args=block_swap_args, compile_args=model["compile_args"])
if gguf_reader is not None: #handle GGUF
load_weights(transformer, patcher.model["sd"], base_dtype=dtype, transformer_load_device=device, patcher=patcher, gguf=True, reader=gguf_reader, block_swap_args=block_swap_args)
load_weights(transformer, patcher.model["sd"], base_dtype=dtype, transformer_load_device=device, patcher=patcher, gguf=True,
reader=gguf_reader, block_swap_args=block_swap_args, compile_args=model["compile_args"])
set_lora_params_gguf(transformer, patcher.patches)
transformer.patched_linear = True
elif len(patcher.patches) != 0: #handle patched linear layers (unmerged loras, fp8 scaled)
+1 -1
View File
@@ -172,7 +172,7 @@ class WanVideoDiffusionForcingSampler:
# Load weights
if not transformer.patched_linear and patcher.model["sd"] is not None and len(patcher.patches) != 0:
transformer = _replace_linear(transformer, dtype, patcher.model["sd"])
transformer = _replace_linear(transformer, dtype, patcher.model["sd"], model["compile_args"])
transformer.patched_linear = True
if patcher.model["sd"] is not None and gguf_reader is None:
load_weights(patcher.model.diffusion_model, patcher.model["sd"], weight_dtype, base_dtype=dtype, transformer_load_device=device, block_swap_args=block_swap_args)