revert this

This commit is contained in:
kijai
2025-07-08 16:23:29 +03:00
parent 948805b6e5
commit ce0bae1839
+3 -4
View File
@@ -83,7 +83,6 @@ class GGUFLinear(nn.Linear):
def forward(self, inputs):
weight = dequantize_without_compile(self.weight)
temp_weight = weight.to(torch.float32)
weight = weight.to(self.compute_dtype)
bias = self.bias.to(self.compute_dtype) if self.bias is not None else None
@@ -92,12 +91,12 @@ class GGUFLinear(nn.Linear):
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(temp_weight.device),
lora_diff[1].flatten(start_dim=1).to(temp_weight.device)
lora_diff[0].flatten(start_dim=1).to(weight.device),
lora_diff[1].flatten(start_dim=1).to(weight.device)
).reshape(weight.shape)
# Apply the patch with its strength
weight = (temp_weight + ((lora_strength * lora_alpha) * patch_diff)).to(self.compute_dtype)
weight = weight + ((lora_strength * lora_alpha) * patch_diff).to(self.compute_dtype)
output = torch.nn.functional.linear(inputs, weight, bias)
return output