Update fast_lora.py

This commit is contained in:
Daniel Han-Chen 2024-03-10 19:09:10 +11:00
commit c192ce3ed4

View file

@ -151,14 +151,16 @@ class LoRA_MLP(torch.autograd.Function):
# dX = matmul_lora(df, upW.t(), upW_quant, upB, upA, upS)
# dX += matmul_lora(de, gateW.t(), gateW_quant, gateB, gateA, gateS)
upW = fast_dequantize(upW.t(), upW_quant)
upW.addmm_(upA.to(dtype), upB.to(dtype), alpha = upS)
dX = torch.matmul(df, upW.t(), out = X)
del upW
dX += df @ upB.to(dtype).t() @ (upS * upA.to(dtype).t())
# dX += df @ upB.to(dtype).t() @ (upS * upA.to(dtype).t())
gateW = fast_dequantize(gateW.t(), gateW_quant)
gateW.addmm_(gateA.to(dtype), gateB.to(dtype), alpha = gateS)
dX += de @ gateW.t()
del gateW
dX += de @ gateB.to(dtype).t() @ (gateS * gateA.to(dtype).t())
# dX += de @ gateB.to(dtype).t() @ (gateS * gateA.to(dtype).t())
# gateW, gateW_quant, gateA, gateB, gateS,
# upW, upW_quant, upA, upB, upS,
@ -306,21 +308,24 @@ class LoRA_QKV(torch.autograd.Function):
# Combine derivatives to find dX
# dQ
QW = fast_dequantize(QW.t(), QW_quant)
QW.addmm_(QA.to(dtype), QA.to(dtype), alpha = QS)
dX = torch.matmul(dQ, QW.t(), out = X)
del QW
dX += (dQ @ QB.to(dtype).t() @ (QS * QA.to(dtype).t()))
# dX += (dQ @ QB.to(dtype).t() @ (QS * QA.to(dtype).t()))
# dK
KW = fast_dequantize(KW.t(), KW_quant)
KW.addmm_(KB.to(dtype), KA.to(dtype), alpha = KS)
dX += dK @ KW.t()
del KW
dX += dK @ KB.to(dtype).t() @ (KS * KA.to(dtype).t())
# dX += dK @ KB.to(dtype).t() @ (KS * KA.to(dtype).t())
# dV
VW = fast_dequantize(VW.t(), VW_quant)
VW.addmm_(VB.to(dtype), VA.to(dtype), alpha = VS)
dX += dV @ VW.t()
del VW
dX += dV @ VB.to(dtype).t() @ (VS * VA.to(dtype).t())
# dX += dV @ VB.to(dtype).t() @ (VS * VA.to(dtype).t())
# QW, QW_quant, QA, QB, QS,
# KW, KW_quant, KA, KB, KS,
@ -406,9 +411,10 @@ class LoRA_W(torch.autograd.Function):
# Get derivative for dX
W = fast_dequantize(W.t(), W_quant)
W.addmm_(A.to(dtype), B.to(dtype), alpha = S)
dX = dY @ W.t()
del W
dX += dY @ B.to(dtype).t() @ (S * A.to(dtype).t())
# dX += dY @ B.to(dtype).t() @ (S * A.to(dtype).t())
# W, W_quant, A, B, S
return dX.view(batch, seq_len, hd), \