From 2cd4b8debd5d2bfdea6da0451404e06cfa75a039 Mon Sep 17 00:00:00 2001 From: Daniel Han Date: Mon, 4 Nov 2024 23:11:03 -0800 Subject: [PATCH] Update rms_layernorm.py --- unsloth/kernels/rms_layernorm.py | 44 ++++++++++++++++---------------- 1 file changed, 22 insertions(+), 22 deletions(-) diff --git a/unsloth/kernels/rms_layernorm.py b/unsloth/kernels/rms_layernorm.py index 7a788f36a7..b6ffa5fee9 100644 --- a/unsloth/kernels/rms_layernorm.py +++ b/unsloth/kernels/rms_layernorm.py @@ -138,7 +138,7 @@ class Fast_RMS_Layernorm(torch.autograd.Function): def forward(ctx, X, W, eps : float, gemma : bool = False): shape = X.shape dim : int = shape[-1] - X : torch.Tensor = X.view(-1, dim) + X = X.view(-1, dim) n_rows : int n_cols : int n_rows, n_cols = X.shape @@ -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 not gemma: + _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 = 16, + ) + 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