From d96f665510569b85d6671dbf04a5245de36c4bdb Mon Sep 17 00:00:00 2001 From: Daniel Han-Chen Date: Fri, 26 Jan 2024 22:21:54 +1100 Subject: [PATCH] Update swiglu.py --- unsloth/kernels/swiglu.py | 4 ++-- 1 file changed, 2 insertions(+), 2 deletions(-) diff --git a/unsloth/kernels/swiglu.py b/unsloth/kernels/swiglu.py index 9948d94102..ec233930a7 100644 --- a/unsloth/kernels/swiglu.py +++ b/unsloth/kernels/swiglu.py @@ -75,8 +75,8 @@ def _DWf_DW_dfg_kernel(DW, e, g, n_elements, BLOCK_SIZE : tl.constexpr,): # dh/dgate = sigmoid(gate)*up + gate*up*sigmoid'(gate) # dh/dgate = sigmoid(gate)*up + gate*up*[sigmoid(gate) * (1 - sigmoid(gate))] # dh/dgate = sigmoid(gate)*up * [1 + gate*(1 - sigmoid(gate))] - # DW_dfg_row = DW_row * se_row * g_row * (1.0 + e_row*(1.0 - se_row)) # 5 FMAs / mults - DW_dfg_row = DW_row * (se_row * g_row + h_row*(1.0 - se_row)) # 4 FMAs / mults + DW_dfg_row = DW_row * se_row * g_row * (1.0 + e_row*(1.0 - se_row)) # 5 FMAs / mults + # DW_dfg_row = DW_row * (se_row * g_row + h_row*(1.0 - se_row)) # 4 FMAs / mults # DW_dfg_row = DW_row * (se_row*(g_row - h_row) + h_row) BREAKS bad accuracy # Store derivatives in buffers