Apply possible LoRA alpha with GGUF too

This commit is contained in:
kijai
2025-07-08 09:27:22 +03:00
parent b355f2b839
commit d16e2aa2ff
+8 -4
View File
@@ -35,6 +35,7 @@ def _replace_with_gguf_linear(model, compute_dtype, state_dict, prefix="", modul
if len(patch) != 0:
lora_diffs = [p[1].weights for p in patch]
lora_strengths = [p[0] for p in patch]
lora_alphas = [p[2] for p in patch]
#print("lora_diff", lora_diff)
@@ -50,7 +51,8 @@ def _replace_with_gguf_linear(model, compute_dtype, state_dict, prefix="", modul
module.bias is not None,
compute_dtype=compute_dtype,
lora_diffs=lora_diffs,
lora_strengths = lora_strengths
lora_strengths = lora_strengths,
lora_alphas = lora_alphas
)
model._modules[name].source_cls = type(module)
# Force requires_grad to False to avoid unexpected errors
@@ -67,12 +69,14 @@ class GGUFLinear(nn.Linear):
compute_dtype=None,
device=None,
lora_diffs=None,
lora_strengths=1.0
lora_strengths=None,
lora_alphas=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_alphas = lora_alphas
def forward(self, inputs):
weight = dequantize_gguf_tensor(self.weight)
@@ -81,7 +85,7 @@ class GGUFLinear(nn.Linear):
if self.lora_diffs is not None:
# Apply all LoRA patches
for lora_diff, lora_strength in zip(self.lora_diffs, self.lora_strengths):
for lora_diff, lora_strength, lora_alpha in zip(self.lora_diffs, self.lora_strengths, self.lora_alphas):
# Calculate the diff for this patch
patch_diff = torch.mm(
lora_diff[0].flatten(start_dim=1).to(weight.device),
@@ -89,7 +93,7 @@ class GGUFLinear(nn.Linear):
).reshape(weight.shape)
# Apply the patch with its strength
weight = weight + patch_diff.to(weight.device, self.compute_dtype) * lora_strength
weight = weight + patch_diff.to(weight.device, self.compute_dtype) * (lora_strength * lora_alpha)
output = torch.nn.functional.linear(inputs, weight, bias)
return output