diff --git a/unsloth/kernels/fast_lora.py b/unsloth/kernels/fast_lora.py index b70e6e4fee..881cbaa14e 100644 --- a/unsloth/kernels/fast_lora.py +++ b/unsloth/kernels/fast_lora.py @@ -128,20 +128,20 @@ 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 = (h.t() @ dY) @ downB.t() + d_downB = 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 = (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() @ (DW_dfg @ gateB.t()) - d_gateB = (gateA.t() @ X.t() @ DW_dfg) + d_gateA = (X.t() @ DW_dfg) @ gateB.t() + d_gateB = gateA.t() @ (X.t() @ DW_dfg) d_gateA *= gateS d_gateB *= gateS @@ -152,13 +152,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 - dX += DW_f @ upB.to(dtype).t() @ (upS * upA.to(dtype).t()) + dX += DW_f @ upB.t() @ (upS * upA.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()) + dX += DW_dfg @ gateB.t() @ (gateS * gateA.t()) # gateW, gateW_quant, gateA, gateB, gateS, # upW, upW_quant, upA, upB, upS, diff --git a/unsloth/kernels/swiglu.py b/unsloth/kernels/swiglu.py index 3106656f32..4e9b7ba2ae 100644 --- a/unsloth/kernels/swiglu.py +++ b/unsloth/kernels/swiglu.py @@ -24,11 +24,11 @@ def _fg_kernel(e, g, h, n_elements, BLOCK_SIZE : tl.constexpr,): offsets = block_idx*BLOCK_SIZE + tl.arange(0, BLOCK_SIZE) mask = offsets < n_elements - e_row = tl.load(e + offsets, mask = mask, other = 0)#.to(tl.float32) + e_row = tl.load(e + offsets, mask = mask, other = 0).to(tl.float32) g_row = tl.load(g + offsets, mask = mask, other = 0)#.to(tl.float32) # f = e * sigmoid(e) - f_row = e_row / (1 + tl.exp(-e_row.to(tl.float32)).to(g_row.dtype)) + f_row = e_row / (1 + tl.exp(-e_row)) f_row = f_row.to(g_row.dtype) # Exact copy from HF # h = f * g h_row = f_row * g_row