diff --git a/unsloth/kernels/layernorm.py b/unsloth/kernels/layernorm.py index 5e2e3af2f8..9e64c3d341 100644 --- a/unsloth/kernels/layernorm.py +++ b/unsloth/kernels/layernorm.py @@ -55,7 +55,9 @@ def layernorm_forward( # (X[0] - mean) == -mean so we need to mask it out XX = tl.where(mask, X_row - mean_X, 0) row_var = tl.sum(XX * XX, axis = 0) / n_cols - inv_var = tl.math.rsqrt(row_var + eps) + # Explicit float32 scalar to ensure correct type promotion on HIP/ROCm + eps_f32 = tl.full((), eps, tl.float32) + inv_var = tl.math.rsqrt(row_var + eps_f32) tl.store(r, inv_var) tl.store(mu, mean_X) output = (XX * inv_var) * W_row + b_row diff --git a/unsloth/kernels/rms_layernorm.py b/unsloth/kernels/rms_layernorm.py index 82e0cd0e9b..74c16c1e63 100644 --- a/unsloth/kernels/rms_layernorm.py +++ b/unsloth/kernels/rms_layernorm.py @@ -49,7 +49,9 @@ def _rms_layernorm_forward( 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) + # Explicit float32 scalar to ensure correct type promotion on HIP/ROCm + eps_f32 = tl.full((), eps, tl.float32) + inv_var = tl.math.rsqrt(row_var + eps_f32) tl.store(r, inv_var) normed = X_row * inv_var normed = normed.to(W_row.dtype) # Exact copy from HF @@ -147,7 +149,9 @@ def _gemma_rms_layernorm_forward( 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) + # Explicit float32 scalar to ensure correct type promotion on HIP/ROCm + eps_f32 = tl.full((), eps, tl.float32) + inv_var = tl.math.rsqrt(row_var + eps_f32) tl.store(r, inv_var) normed = X_row * inv_var output = normed * (W_row + 1.0)