Further GGUF+LoRA fixes
This commit is contained in:
+25
-30
@@ -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
|
||||
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
|
||||
Reference in New Issue
Block a user