diff --git a/unsloth/kernels/rms_layernorm.py b/unsloth/kernels/rms_layernorm.py index 4079921f91..3176d4e358 100644 --- a/unsloth/kernels/rms_layernorm.py +++ b/unsloth/kernels/rms_layernorm.py @@ -142,14 +142,14 @@ class Fast_RMS_Layernorm(torch.autograd.Function): n_rows : int n_cols : int n_rows, n_cols = X.shape + BLOCK_SIZE : int + num_warps : int + BLOCK_SIZE, num_warps = calculate_settings(n_cols) Y = torch.empty((n_rows, n_cols), dtype = X.dtype, device = "cuda:0") r = torch.empty(n_rows, dtype = torch.float32, device = "cuda:0") - BLOCK_SIZE : int - num_warps : int if not gemma: - BLOCK_SIZE, num_warps = calculate_settings(n_cols) _rms_layernorm_forward[(n_rows,)]( Y, Y.stride(0), X, X.stride(0), @@ -158,7 +158,7 @@ class Fast_RMS_Layernorm(torch.autograd.Function): n_cols = int(n_cols), eps = float(eps), BLOCK_SIZE = BLOCK_SIZE, - num_warps = num_warps, + num_warps = 16, ) else: _gemma_rms_layernorm_forward[(n_rows,)]( @@ -168,7 +168,7 @@ class Fast_RMS_Layernorm(torch.autograd.Function): r, r.stride(0), n_cols, eps, BLOCK_SIZE = BLOCK_SIZE, - num_warps = num_warps, + num_warps = 16, ) ctx.eps = eps ctx.BLOCK_SIZE = BLOCK_SIZE