From 15fdce80b9a5f983a1cb1005a20512865b3e2fdd Mon Sep 17 00:00:00 2001 From: Daniel Han Date: Mon, 4 Nov 2024 23:08:05 -0800 Subject: [PATCH] Update rms_layernorm.py --- unsloth/kernels/rms_layernorm.py | 42 ++++++++++++++++---------------- 1 file changed, 21 insertions(+), 21 deletions(-) diff --git a/unsloth/kernels/rms_layernorm.py b/unsloth/kernels/rms_layernorm.py index 4074d3a502..7a788f36a7 100644 --- a/unsloth/kernels/rms_layernorm.py +++ b/unsloth/kernels/rms_layernorm.py @@ -149,27 +149,27 @@ class Fast_RMS_Layernorm(torch.autograd.Function): Y = torch.empty((n_rows, n_cols), dtype = X.dtype, device = "cuda:0") r = torch.empty(n_rows, dtype = torch.float32, device = "cuda:0") - if gemma == False: - _rms_layernorm_forward[(n_rows,)]( - Y, Y.stride(0), - X, X.stride(0), - W, W.stride(0), - r, r.stride(0), - n_cols = int(n_cols), - eps = float(eps), - BLOCK_SIZE = BLOCK_SIZE, - num_warps = num_warps, - ) - else: - _gemma_rms_layernorm_forward[(n_rows,)]( - Y, Y.stride(0), - X, X.stride(0), - W, W.stride(0), - r, r.stride(0), - n_cols, eps, - BLOCK_SIZE = BLOCK_SIZE, - num_warps = num_warps, - ) + # if gemma == False: + _rms_layernorm_forward[(n_rows,)]( + Y, Y.stride(0), + X, X.stride(0), + W, W.stride(0), + r, r.stride(0), + n_cols = int(n_cols), + eps = float(eps), + BLOCK_SIZE = BLOCK_SIZE, + num_warps = num_warps, + ) + # else: + # _gemma_rms_layernorm_forward[(n_rows,)]( + # Y, Y.stride(0), + # X, X.stride(0), + # W, W.stride(0), + # r, r.stride(0), + # n_cols, eps, + # BLOCK_SIZE = BLOCK_SIZE, + # num_warps = num_warps, + # ) ctx.eps = eps ctx.BLOCK_SIZE = BLOCK_SIZE ctx.num_warps = num_warps