Update fast_lora.py

This commit is contained in:
Daniel Han-Chen 2024-01-27 03:23:50 +11:00
commit 421ed33b42

View file

@ -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()))