From 421ed33b42ccc040badbd7193bc8da4aa383448d Mon Sep 17 00:00:00 2001 From: Daniel Han-Chen Date: Sat, 27 Jan 2024 03:23:50 +1100 Subject: [PATCH] Update fast_lora.py --- unsloth/kernels/fast_lora.py | 26 +++++++++++++------------- 1 file changed, 13 insertions(+), 13 deletions(-) diff --git a/unsloth/kernels/fast_lora.py b/unsloth/kernels/fast_lora.py index 82bd4c0c19..aaea063682 100644 --- a/unsloth/kernels/fast_lora.py +++ b/unsloth/kernels/fast_lora.py @@ -139,29 +139,29 @@ class LoRA_MLP(torch.autograd.Function): # h, DW_f, DW_dfg = DW, e, g # Down projection LoRA weights - d_downA = h.t() @ (dY @ downB.t()) - d_downB = (downA.t() @ h.t()) @ dY - d_downA *= downS - d_downB *= downS + d_downA = h.t() @ (dY @ (downS * downB).t()) + d_downB = ((downS * downA).t() @ h.t()) @ dY + # d_downA *= downS + # d_downB *= downS # Up projection LoRA weights - d_upA = X.t() @ (DW_f @ upB.t()) - d_upB = (upA.t() @ X.t()) @ DW_f - d_upA *= upS - d_upB *= upS + d_upA = X.t() @ (DW_f @ (upS * upB).t()) + d_upB = ((upS * upA).t() @ X.t()) @ DW_f + # d_upA *= upS + # d_upB *= upS # Gate projection LoRA weights - d_gateA = X.t() @ (DW_dfg @ gateB.t()) - d_gateB = (gateA.t() @ X.t()) @ DW_dfg - d_gateA *= gateS - d_gateB *= gateS + d_gateA = X.t() @ (DW_dfg @ (gateS * gateB).t()) + d_gateB = ((gateS * gateA).t() @ X.t()) @ DW_dfg + # d_gateA *= gateS + # d_gateB *= gateS # Final derivatives to backpropagate backwards. # See our blogpost for more details. # (D @ W.T * f) @ U.T upW = fast_dequantize(upW.t(), upW_quant) # (D @ W.T * f) @ (U.T + B.T @ A.T) - dX = torch.matmul(DW_f, upW.t()) + dX = torch.matmul(DW_f, upW.t(), out = X) del upW dX += (DW_f @ upB.to(dtype).t() @ (upS * upA.to(dtype).t()))