Apply possible LoRA alpha with GGUF too
This commit is contained in:
+8
-4
@@ -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
|
||||
Reference in New Issue
Block a user