From 47babc780af9a6c724e43de422bec6652af4e986 Mon Sep 17 00:00:00 2001 From: Daniel Han-Chen Date: Sat, 27 Jan 2024 04:16:38 +1100 Subject: [PATCH] Update fast_lora.py --- unsloth/kernels/fast_lora.py | 159 ++++++----------------------------- 1 file changed, 26 insertions(+), 133 deletions(-) diff --git a/unsloth/kernels/fast_lora.py b/unsloth/kernels/fast_lora.py index 550539d927..b487ff95e2 100644 --- a/unsloth/kernels/fast_lora.py +++ b/unsloth/kernels/fast_lora.py @@ -90,8 +90,7 @@ 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) i = matmul_lora(h, downW, downW_quant, downA, downB, downS) ctx.custom_saved_tensors = ( @@ -122,126 +121,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 = 1 / (1 + torch.exp(-e.float())) - f = torch.nn.functional.silu(e) - h = f * g - DW_f = (DW * f) - DW_dfg = (DW * g).float() * se * (1.0 + e.float() * (1.0 - se)) - DW_dfg = DW_dfg.to(dtype) - # 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 - - # Down projection LoRA weights - 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 @ (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 @ (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(), out = X) - del upW - 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 @ gateB.to(dtype).t() @ (gateS * gateA.to(dtype).t()) - dX += DW_dfg @ gateW.t() - del gateW - - # gateW, gateW_quant, gateA, gateB, gateS, - # upW, upW_quant, upA, upB, upS, - # downW, downW_quant, downA, downB, downS, - return dX.view(batch, seq_len, hd), \ - None, None, d_gateA.t(), d_gateB.t(), None, \ - None, None, d_upA.t(), d_upB.t(), None, \ - None, None, d_downA.t(), d_downB.t(), None, - pass -pass - - -class LoRA_MLP_New(torch.autograd.Function): - @classmethod - @torch.cuda.amp.custom_fwd - def forward(cls, ctx, X : torch.Tensor, - gateW, gateW_quant, gateA, gateB, gateS, - upW, upW_quant, upA, upB, upS, - downW, downW_quant, downA, downB, downS): - dtype = X.dtype - - e = matmul_lora(X, gateW, gateW_quant, gateA, gateB, gateS) - g = matmul_lora(X, upW, upW_quant, upA, upB, upS) - f = torch.nn.functional.silu(e) - h = f * g - i = matmul_lora(h, downW, downW_quant, downA, downB, downS) - - ctx.custom_saved_tensors = ( - gateW, gateW_quant, gateS, - upW, upW_quant, upS, - downW, downW_quant, downS, - ) - ctx.save_for_backward(gateA, gateB, upA, upB, downA, downB, - X, e, g, f, h, i) - return i - pass - - def _silu_backward(dy, X): - # https://github.com/pytorch/pytorch/blob/563b065f5a4b4055fa6b025c2514b566d5fd9439/aten/src/ATen/native/Activation.cpp#L483 - sigm = 1 / (1 + torch.exp(-X.float())) - return (dy.float() * sigm * (1 + X.float() * (1 - sigm))).to(X.dtype) - pass - - @classmethod - @torch.cuda.amp.custom_bwd - def backward(cls, ctx, dY : torch.Tensor): - gateW, gateW_quant, gateS, upW, upW_quant, upS, downW, downW_quant, downS, = \ - ctx.custom_saved_tensors - gateA, gateB, upA, upB, downA, downB, \ - X, e, g, f, h, i = ctx.saved_tensors - - gateA, gateB, upA, upB, downA, downB = \ - gateA.t(), gateB.t(), upA.t(), upB.t(), downA.t(), downB.t() - - batch, seq_len, hd = X.shape - dY = dY.view(-1, dY.shape[-1]) - X = X .view(-1, X .shape[-1]) - e = e .view(-1, e .shape[-1]) - g = g .view(-1, g .shape[-1]) - f = f .view(-1, f .shape[-1]) - h = h .view(-1, h .shape[-1]) - i = i .view(-1, i .shape[-1]) - dtype = X.dtype - - DW = matmul_lora(dY, downW.t(), downW_quant, downB, downA, downS) - df = DW * f # 88us - dg = DW * g # 88us - de = cls._silu_backward(dg, e) # 90us - - dX = matmul_lora(df, upW.t(), upW_quant, upB, upA, upS) - dX += matmul_lora(de, gateW.t(), gateW_quant, gateB, gateA, gateS) + 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()) @@ -250,17 +134,32 @@ class LoRA_MLP_New(torch.autograd.Function): d_downB *= downS # Up projection LoRA weights - d_upA = X.t() @ (df @ upB.t()) - d_upB = (upA.t() @ X.t()) @ df + d_upA = X.t() @ (DW_f @ upB.t()) + d_upB = (upA.t() @ X.t()) @ DW_f d_upA *= upS d_upB *= upS # Gate projection LoRA weights - d_gateA = X.t() @ (dg @ gateB.t()) - d_gateB = (gateA.t() @ X.t()) @ dg + d_gateA = X.t() @ (DW_dfg @ gateB.t()) + d_gateB = (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(), out = X) + del upW + 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) + dX += DW_dfg @ gateW.t() + del gateW + dX += DW_dfg @ gateB.to(dtype).t() @ (gateS * gateA.to(dtype).t()) + # gateW, gateW_quant, gateA, gateB, gateS, # upW, upW_quant, upA, upB, upS, # downW, downW_quant, downA, downB, downS, @@ -271,18 +170,12 @@ class LoRA_MLP_New(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) - out = LoRA_MLP_New.apply(X, + out = LoRA_MLP.apply(X, gateW, gateW_quant, gateA, gateB, gateS, upW, upW_quant, upA, upB, upS, downW, downW_quant, downA, downB, downS)