Update fast_lora.py
This commit is contained in:
parent
ae0d9380c1
commit
f8aa20d2db
1 changed files with 11 additions and 27 deletions
|
|
@ -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)
|
||||
|
|
|
|||
Loading…
Add table
Add a link
Reference in a new issue