From 794a38cc7ea37ae2b87a37118a1e1817c8c3b292 Mon Sep 17 00:00:00 2001 From: kijai <40791699+kijai@users.noreply.github.com> Date: Fri, 25 Jul 2025 00:40:48 +0300 Subject: [PATCH] Further GGUF+LoRA fixes --- gguf/gguf.py | 55 ++++++++++++++++++++++++---------------------------- 1 file changed, 25 insertions(+), 30 deletions(-) diff --git a/gguf/gguf.py b/gguf/gguf.py index b1d72b3..e1d9655 100644 --- a/gguf/gguf.py +++ b/gguf/gguf.py @@ -8,10 +8,6 @@ if is_accelerate_available(): import accelerate from accelerate import init_empty_weights -@torch.compiler.disable() -def dequantize_without_compile(tensor): - return dequantize_gguf_tensor(tensor) - #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 _should_convert_to_gguf(state_dict, prefix): @@ -52,11 +48,12 @@ def _replace_with_gguf_linear(model, compute_dtype, state_dict, prefix="", modul def set_lora_params(module, patches, module_prefix=""): # Recursively set lora_diffs and lora_strengths for all GGUFLinear layers for name, child in module.named_children(): - child_prefix = (f"{module_prefix}{name}.").replace("_orig_mod.", "") + child_prefix = (f"{module_prefix}{name}.") set_lora_params(child, patches, child_prefix) if isinstance(module, GGUFLinear): key = f"diffusion_model.{module_prefix}weight" patch = patches.get(key, []) + #print(f"Processing LoRA patches for {key}: {len(patch)} patches found") if len(patch) != 0: lora_diffs = [] for p in patch: @@ -70,8 +67,8 @@ def set_lora_params(module, patches, module_prefix=""): else: continue lora_strengths = [p[0] for p in patch] - module.lora_diffs = lora_diffs - module.lora_strengths = lora_strengths + module.lora = (lora_diffs, lora_strengths) + class GGUFLinear(nn.Linear): def __init__( @@ -81,36 +78,34 @@ class GGUFLinear(nn.Linear): bias=False, compute_dtype=None, device=None, - lora_diffs=None, - lora_strengths=None, ) -> None: super().__init__(in_features, out_features, bias, device) self.compute_dtype = compute_dtype - self.lora_diffs = lora_diffs - self.lora_strengths = lora_strengths + self.lora = None def forward(self, inputs): - weight = dequantize_without_compile(self.weight) + weight = self.dequantize_without_compile() weight = weight.to(self.compute_dtype) bias = self.bias.to(self.compute_dtype) if self.bias is not None else None - if self.lora_diffs is not None: - # Apply all LoRA patches - for lora_diff, lora_strength in zip(self.lora_diffs, self.lora_strengths): - # Calculate the diff for this patch - 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) - - if lora_diff[2] is not None: - alpha = lora_diff[2] / lora_diff[1].shape[0] - else: - alpha = 1.0 - - # Apply the patch with its strength - scale = lora_strength * alpha - weight.add_(patch_diff, alpha=scale).to(self.compute_dtype) + if hasattr(self, "lora"): + weight = self.apply_lora(weight).to(self.compute_dtype) output = torch.nn.functional.linear(inputs, weight, bias) - return output \ No newline at end of file + return output + + @torch.compiler.disable() + def dequantize_without_compile(self): + return dequantize_gguf_tensor(self.weight) + + @torch.compiler.disable() + def apply_lora(self, weight): + for lora_diff, lora_strength in zip(self.lora[0], self.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 \ No newline at end of file