diff --git a/unsloth/kernels/fast_lora.py b/unsloth/kernels/fast_lora.py index e24fa48dcc..52cbf47488 100644 --- a/unsloth/kernels/fast_lora.py +++ b/unsloth/kernels/fast_lora.py @@ -123,16 +123,16 @@ class LoRA_MLP(torch.autograd.Function): g = g .view(-1, g .shape[-1]) dtype = X.dtype - DW = matmul_lora(dY, downW.t(), downW_quant, downB, downA, downS) - se = torch.nn.functional.sigmoid(e) - 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 + # se = torch.nn.functional.sigmoid(e) + # 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 @ downB.t())