From 74fc5caa60c712757165294134cd1b5b49e7f722 Mon Sep 17 00:00:00 2001 From: Daniel Han-Chen Date: Sun, 10 Mar 2024 16:20:20 +1100 Subject: [PATCH] Accuracy --- unsloth/kernels/rms_layernorm.py | 2 +- unsloth/kernels/swiglu.py | 6 +++--- 2 files changed, 4 insertions(+), 4 deletions(-) diff --git a/unsloth/kernels/rms_layernorm.py b/unsloth/kernels/rms_layernorm.py index 4db89b7816..2adedff663 100644 --- a/unsloth/kernels/rms_layernorm.py +++ b/unsloth/kernels/rms_layernorm.py @@ -41,7 +41,7 @@ def _rms_layernorm_forward( r += row_idx * r_row_stride X_row = tl.load(X + col_offsets, mask = mask, other = 0).to(tl.float32) - W_row = tl.load(W + col_offsets, mask = mask, other = 0)#.to(tl.float32) + W_row = tl.load(W + col_offsets, mask = mask, other = 0).to(tl.float32) row_var = tl.sum(X_row * X_row, axis = 0) / n_cols inv_var = tl.math.rsqrt(row_var + eps) diff --git a/unsloth/kernels/swiglu.py b/unsloth/kernels/swiglu.py index ff6b162680..4f4fbe9ee9 100644 --- a/unsloth/kernels/swiglu.py +++ b/unsloth/kernels/swiglu.py @@ -25,7 +25,7 @@ def _fg_kernel(e, g, h, n_elements, BLOCK_SIZE : tl.constexpr,): mask = offsets < n_elements 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) + g_row = tl.load(g + offsets, mask = mask, other = 0).to(tl.float32) # f = e * sigmoid(e) f_row = e_row * tl.sigmoid(e_row) # e_row / (1 + tl.exp(-e_row)) @@ -63,9 +63,9 @@ def _DWf_DW_dfg_kernel(DW, e, g, n_elements, BLOCK_SIZE : tl.constexpr,): offsets = block_idx*BLOCK_SIZE + tl.arange(0, BLOCK_SIZE) mask = offsets < n_elements - DW_row = tl.load(DW + offsets, mask = mask, other = 0)#.to(tl.float32) + DW_row = tl.load(DW + 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) + g_row = tl.load(g + offsets, mask = mask, other = 0).to(tl.float32) # e = e.float() # se = 1.0 / (1.0 + torch.exp(-e))