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:
+41
-24
@@ -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
@@ -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)
|
||||
|
||||
@@ -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
@@ -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
@@ -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)
|
||||
|
||||
Reference in New Issue
Block a user