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 <danielhanchen@gmail.com>
This commit is contained in:
parent
dc75d00d14
commit
e4daae62d9
2 changed files with 9 additions and 3 deletions
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
Loading…
Add table
Add a link
Reference in a new issue