From 17795e4f141fb7b90d652d8f2d423d3c8bd89bc1 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?=E9=87=91=E9=BB=84=E8=89=B2=E8=91=A1=E8=90=84=E7=90=83?= =?UTF-8?q?=E5=90=9B=E5=90=9B?= Date: Sun, 1 Mar 2026 15:59:22 +0800 Subject: [PATCH] fix(Triton): ensure float32 eps in RMS LayerNorm rsqrt for HIP/ROCm (#4110) * fix(Triton): ensure float32 eps in RMS LayerNorm rsqrt for HIP/ROCm On HIP (AMD ROCm), Triton constexpr eps may not promote to float32 in rsqrt, causing numerical instability (NaN/Inf) on RDNA GPUs (gfx1100, gfx1151 Strix Halo, etc.). Use tl.full((), eps, tl.float32) to explicitly create a float32 scalar before adding to row_var in rsqrt. Applied to both standard and Gemma RMS LayerNorm forward kernels. Tested on W7900 (gfx1100): full test suite passed (dim 512-2048, bf16/fp16, various seqlen). Related: #3385, #3588 * Apply same float32 eps fix to layernorm.py for PR #4110 layernorm.py has the identical tl.constexpr eps pattern in layernorm_forward that can misfire on HIP/ROCm. Apply the same tl.full((), eps, tl.float32) fix for consistency. Both testing_suite_layernorm (standard LayerNorm) and testing_suite_layernorm (RMS LayerNorm) pass on NVIDIA after this change. --------- Co-authored-by: Daniel Han --- unsloth/kernels/layernorm.py | 4 +++- unsloth/kernels/rms_layernorm.py | 8 ++++++-- 2 files changed, 9 insertions(+), 3 deletions(-) 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)