diff --git a/unsloth/kernels/fast_lora.py b/unsloth/kernels/fast_lora.py index 6f8ea2a68e..7f97fea5c7 100644 --- a/unsloth/kernels/fast_lora.py +++ b/unsloth/kernels/fast_lora.py @@ -17,6 +17,7 @@ from .utils import fast_dequantize, QUANT_STATE, get_lora_parameters from .swiglu import swiglu_fg_kernel, swiglu_DWf_DW_dfg_kernel + def matmul_lora(X, W, W_quant, A, B, s, out = None): dtype = X.dtype W = fast_dequantize(W.t(), W_quant) @@ -90,8 +91,8 @@ class LoRA_MLP(torch.autograd.Function): e = matmul_lora(X, gateW, gateW_quant, gateA, gateB, gateS) g = matmul_lora(X, upW, upW_quant, upA, upB, upS) - h = torch.nn.functional.silu(e) * g - # h = swiglu_fg_kernel(e, g) + h = swiglu_fg_kernel(e, g) + # h = torch.nn.functional.silu(e) * g i = matmul_lora(h, downW, downW_quant, downA, downB, downS) ctx.custom_saved_tensors = ( @@ -122,20 +123,11 @@ class LoRA_MLP(torch.autograd.Function): g = g .view(-1, g .shape[-1]) dtype = X.dtype + # DW_f = (D @ W.T * f) + # DW_dfg = (D @ W.T * df * g) DW = matmul_lora(dY, downW.t(), downW_quant, downB, downA, downS) - se = torch.nn.functional.sigmoid(e) - f = torch.nn.functional.silu(e) - h = f * g - DW_f = DW * f - DW_dfg = DW * se * g * (1 + f * (1 - se)) - # f = e * se - # h = f * g - # df = se * (1 - f) + f - # DW_f = DW * f - # DW_dfg = DW * df * g - # DW = matmul_lora(dY, downW.t(), downW_quant, downB, downA, downS) - # DW, e, g = swiglu_DWf_DW_dfg_kernel(DW, e, g) - # h, DW_f, DW_dfg = DW, e, g + DW, e, g = swiglu_DWf_DW_dfg_kernel(DW, e, g) + h, DW_f, DW_dfg = DW, e, g # Down projection LoRA weights d_downA = h.t() @ (dY @ downB.t()) @@ -162,15 +154,13 @@ class LoRA_MLP(torch.autograd.Function): # (D @ W.T * f) @ (U.T + B.T @ A.T) dX = torch.matmul(DW_f, upW.t(), out = X) del upW - new_dX = upS * (DW_f @ upB.to(dtype).t() @ (upA.to(dtype).t())) + dX += DW_f @ upB.to(dtype).t() @ (upS * upA.to(dtype).t()) # And add the derivative for the gate projection gateW = fast_dequantize(gateW.t(), gateW_quant) - new_dX2 = DW_dfg @ gateW.t() - # dX += DW_dfg @ gateW.t() + dX += DW_dfg @ gateW.t() del gateW - new_dX2 += gateS * (DW_dfg @ gateB.to(dtype).t() @ (gateA.to(dtype).t())) - dX += (new_dX + new_dX2) + dX += DW_dfg @ gateB.to(dtype).t() @ (gateS * gateA.to(dtype).t()) # gateW, gateW_quant, gateA, gateB, gateS, # upW, upW_quant, upA, upB, upS, @@ -182,14 +172,8 @@ class LoRA_MLP(torch.autograd.Function): pass pass -from transformers.models.llama.modeling_llama import logger + def apply_lora_mlp(self, X): - logger.warning_once("Hello!2") - # gate = self.gate_proj(X) - # up = self. up_proj(X) - # h = torch.nn.functional.silu(gate) * up - # down = self.down_proj(h) - # return down gateW, gateW_quant, gateA, gateB, gateS = get_lora_parameters(self.gate_proj) upW, upW_quant, upA, upB, upS = get_lora_parameters(self. up_proj) downW, downW_quant, downA, downB, downS = get_lora_parameters(self.down_proj)