diff --git a/aie_kernels/aie2/rms_norm.cc b/aie_kernels/aie2/rms_norm.cc index 7c9d7c2b..4aa41e7a 100644 --- a/aie_kernels/aie2/rms_norm.cc +++ b/aie_kernels/aie2/rms_norm.cc @@ -9,10 +9,9 @@ #include template -void rms_norm_general(const T *restrict input, const T *restrict input2, T *restrict output, int32_t cols) +void rms_norm_general(const T *restrict input, const T *restrict input2, T *restrict output, int32_t cols, float epsilon) { event0(); - constexpr float epsilon = 1e-5f; ::aie::vector add_res = ::aie::zeros(); int vector_chunks = cols / N; @@ -35,19 +34,22 @@ void rms_norm_general(const T *restrict input, const T *restrict input2, T *rest float rms = sum_sq / cols + epsilon; float inv_rms = invsqrt(rms); - ::aie::vector inv_rms_v = ::aie::broadcast(static_cast(inv_rms)); + // Normalize in f32 and round once. The previous code cast inv_rms to bf16 and did a + // bf16*bf16 multiply, a coherent per-norm scale error that accumulated over a deep + // residual stream and flipped near-tie argmaxes. Mirrors layer_norm.cc's accfloat path. + ::aie::accum inv_rms_v; + inv_rms_v.from_vector(::aie::broadcast(inv_rms), 0); for (int i = 0; i < vector_chunks; i++) { - ::aie::vector reg_a = ::aie::load_v(input + i * N); - ::aie::vector norm_v = ::aie::mul(reg_a, inv_rms_v); - ::aie::vector out_v; + ::aie::accum reg_a; + reg_a.from_vector(::aie::load_v(input + i * N), 0); + reg_a = ::aie::mul(reg_a.template to_vector(), inv_rms_v.template to_vector()); if (input2) { - ::aie::vector reg_b = ::aie::load_v(input2 + i * N); - out_v = ::aie::mul(norm_v, reg_b); - } else { - out_v = norm_v; + ::aie::accum reg_b; + reg_b.from_vector(::aie::load_v(input2 + i * N), 0); + reg_a = ::aie::mul(reg_a.template to_vector(), reg_b.template to_vector()); } - ::aie::store_v(output + i * N, out_v); + ::aie::store_v(output + i * N, reg_a.template to_vector()); } if (remaining > 0) { @@ -67,13 +69,15 @@ void rms_norm_general(const T *restrict input, const T *restrict input2, T *rest } extern "C" { -void rms_norm_bf16_vector(bfloat16 *input, bfloat16 *output, int32_t size) +void rms_norm_bf16_vector(bfloat16 *input, bfloat16 *output, int32_t size, float epsilon) { - rms_norm_general(input, nullptr, output, size); + ::aie::set_rounding(aie::rounding_mode::conv_even); // round-to-nearest-even; do not inherit a floor rounding mode from a prior kernel + rms_norm_general(input, nullptr, output, size, epsilon); } -void weighted_rms_norm(bfloat16 *a_in, bfloat16 *b_in, bfloat16 *c_out, int32_t size) +void weighted_rms_norm(bfloat16 *a_in, bfloat16 *b_in, bfloat16 *c_out, int32_t size, float epsilon) { - rms_norm_general(a_in, b_in, c_out, size); + ::aie::set_rounding(aie::rounding_mode::conv_even); // round-to-nearest-even; do not inherit a floor rounding mode from a prior kernel + rms_norm_general(a_in, b_in, c_out, size, epsilon); } } diff --git a/aie_kernels/aie2p/rms_norm.cc b/aie_kernels/aie2p/rms_norm.cc index 1a709309..0cfcc3a6 100644 --- a/aie_kernels/aie2p/rms_norm.cc +++ b/aie_kernels/aie2p/rms_norm.cc @@ -7,10 +7,9 @@ #include template -void rms_norm_general(const T *restrict input, const T *restrict input2, T *restrict output, int32_t cols) +void rms_norm_general(const T *restrict input, const T *restrict input2, T *restrict output, int32_t cols, float epsilon) { event0(); - constexpr float epsilon = 1e-5f; ::aie::vector add_res = ::aie::zeros(); int vector_chunks = cols / N; @@ -33,19 +32,22 @@ void rms_norm_general(const T *restrict input, const T *restrict input2, T *rest float rms = sum_sq / cols + epsilon; float inv_rms = aie::invsqrt(rms); - ::aie::vector inv_rms_v = ::aie::broadcast(static_cast(inv_rms)); + // Normalize in f32 and round once. The previous code cast inv_rms to bf16 and did a + // bf16*bf16 multiply, a coherent per-norm scale error that accumulated over a deep + // residual stream and flipped near-tie argmaxes. Mirrors layer_norm.cc's accfloat path. + ::aie::accum inv_rms_v; + inv_rms_v.from_vector(::aie::broadcast(inv_rms), 0); for (int i = 0; i < vector_chunks; i++) { - ::aie::vector reg_a = ::aie::load_v(input + i * N); - ::aie::vector norm_v = ::aie::mul(reg_a, inv_rms_v); - ::aie::vector out_v; + ::aie::accum reg_a; + reg_a.from_vector(::aie::load_v(input + i * N), 0); + reg_a = ::aie::mul(reg_a.template to_vector(), inv_rms_v.template to_vector()); if (input2) { - ::aie::vector reg_b = ::aie::load_v(input2 + i * N); - out_v = ::aie::mul(norm_v, reg_b); - } else { - out_v = norm_v; + ::aie::accum reg_b; + reg_b.from_vector(::aie::load_v(input2 + i * N), 0); + reg_a = ::aie::mul(reg_a.template to_vector(), reg_b.template to_vector()); } - ::aie::store_v(output + i * N, out_v); + ::aie::store_v(output + i * N, reg_a.template to_vector()); } if (remaining > 0) { @@ -65,13 +67,15 @@ void rms_norm_general(const T *restrict input, const T *restrict input2, T *rest } extern "C" { -void rms_norm_bf16_vector(bfloat16 *input, bfloat16 *output, int32_t size) +void rms_norm_bf16_vector(bfloat16 *input, bfloat16 *output, int32_t size, float epsilon) { - rms_norm_general(input, nullptr, output, size); + ::aie::set_rounding(aie::rounding_mode::conv_even); // round-to-nearest-even; do not inherit a floor rounding mode from a prior kernel + rms_norm_general(input, nullptr, output, size, epsilon); } -void weighted_rms_norm(bfloat16 *a_in, bfloat16 *b_in, bfloat16 *c_out, int32_t size) +void weighted_rms_norm(bfloat16 *a_in, bfloat16 *b_in, bfloat16 *c_out, int32_t size, float epsilon) { - rms_norm_general(a_in, b_in, c_out, size); + ::aie::set_rounding(aie::rounding_mode::conv_even); // round-to-nearest-even; do not inherit a floor rounding mode from a prior kernel + rms_norm_general(a_in, b_in, c_out, size, epsilon); } } diff --git a/iron/operators/rms_norm/design.py b/iron/operators/rms_norm/design.py index af96cac9..d2183f62 100644 --- a/iron/operators/rms_norm/design.py +++ b/iron/operators/rms_norm/design.py @@ -17,6 +17,7 @@ def my_rms_norm( num_channels, tile_size, trace_size, + epsilon=1e-5, ): per_tile_elements = 8192 if tile_size > 8192 else tile_size total_cores = num_columns * num_channels @@ -49,7 +50,7 @@ def my_rms_norm( # AIE Core Function declaration rms_norm_kernel = Kernel( - "rms_norm_bf16_vector", "rms_norm.o", [tile_ty, tile_ty, np.int32] + "rms_norm_bf16_vector", "rms_norm.o", [tile_ty, tile_ty, np.int32, np.float32] ) # Define a task that will run on a compute tile @@ -58,7 +59,7 @@ def core_body(of_in1, of_out, rms_norm_kernel): for _ in range_(N_div_n): elem_in1 = of_in1.acquire(1) elem_out = of_out.acquire(1) - rms_norm_kernel(elem_in1, elem_out, per_tile_elements) + rms_norm_kernel(elem_in1, elem_out, per_tile_elements, epsilon) of_in1.release(1) of_out.release(1) diff --git a/iron/operators/rms_norm/design_weighted.py b/iron/operators/rms_norm/design_weighted.py index c0d77fe3..e5333feb 100644 --- a/iron/operators/rms_norm/design_weighted.py +++ b/iron/operators/rms_norm/design_weighted.py @@ -17,6 +17,7 @@ def my_weighted_rms_norm( num_channels, weight_length, trace_size, + epsilon=1e-5, func_prefix="", ): per_tile_elements = weight_length @@ -63,7 +64,7 @@ def my_weighted_rms_norm( rms_norm_kernel = Kernel( f"{func_prefix}rms_norm_bf16_vector", f"{func_prefix}rms_norm.o", - [tile_ty, tile_ty, np.int32], + [tile_ty, tile_ty, np.int32, np.float32], ) eltwise_mul_kernel = Kernel( f"{func_prefix}eltwise_mul_bf16_vector", @@ -77,7 +78,7 @@ def core_body_norm(of_in1, of_out1, rms_norm): for _ in range_(N_div_n): elem_in1 = of_in1.acquire(1) elem_out = of_out1.acquire(1) - rms_norm(elem_in1, elem_out, per_tile_elements) + rms_norm(elem_in1, elem_out, per_tile_elements, epsilon) of_in1.release(1) of_out1.release(1) diff --git a/iron/operators/rms_norm/op.py b/iron/operators/rms_norm/op.py index 1ee6c97d..c103f4c5 100644 --- a/iron/operators/rms_norm/op.py +++ b/iron/operators/rms_norm/op.py @@ -26,15 +26,16 @@ class RMSNorm(MLIROperator): num_channels: int tile_size: int weighted: bool = False + epsilon: float = 1e-5 # RMSNorm eps; Llama 1e-5 (default), Gemma 1e-6 context: object = field(default=None, repr=False) _name_aliases: ClassVar[Dict[str, str]] = { **MLIROperator._name_aliases, "weighted": "w", + "epsilon": "eps", } def __post_init__(self): - # Note: epsilon is hardcoded to 1e-5 in the AIE kernel and cannot be changed at runtime. dev = aie_utils.get_current_device() shim_dma_limit = get_shim_dma_limit(dev) @@ -87,6 +88,7 @@ def get_mlir_artifact(self): self.num_channels, self.tile_size, 0, # trace_size + self.epsilon, ), ), ) @@ -129,4 +131,4 @@ def reference(self, x, w=None): """CPU reference: row-wise RMS normalization, optionally weighted.""" from iron.operators.rms_norm.reference import reference - return reference(x, w=w, weighted=self.weighted) + return reference(x, w=w, weighted=self.weighted, eps=self.epsilon) diff --git a/iron/operators/rms_norm/reference.py b/iron/operators/rms_norm/reference.py index 55ff231f..184ed7da 100644 --- a/iron/operators/rms_norm/reference.py +++ b/iron/operators/rms_norm/reference.py @@ -5,25 +5,28 @@ from iron.common.test_utils import torch_dtype_map -def reference(x, w=None, weighted=False): - """CPU reference: row-wise RMS normalization, optionally weighted (ground truth).""" - rms = torch.sqrt(torch.mean(x**2, dim=-1, keepdim=True)) - out = x / (rms + 1e-5) +def reference(x, w=None, weighted=False, eps=1e-5): + """CPU reference: row-wise RMS normalization, optionally weighted (ground truth). + + Matches the AIE kernel: normalize by 1/sqrt(mean(x^2) + eps). + """ + rms = torch.sqrt(torch.mean(x**2, dim=-1, keepdim=True) + eps) + out = x / rms if weighted: out = out * w return out def generate_golden_reference( - rows: int, cols: int, dtype="bf16", seed=42, weighted=False + rows: int, cols: int, dtype="bf16", seed=42, weighted=False, eps=1e-5 ): torch.manual_seed(seed) val_range = 4 input_tensor = torch.rand(rows, cols, dtype=torch_dtype_map[dtype]) * val_range if weighted: weights = torch.rand(cols, dtype=torch_dtype_map[dtype]) * val_range - output_tensor = reference(input_tensor, weights, weighted=True) + output_tensor = reference(input_tensor, weights, weighted=True, eps=eps) return {"input": input_tensor, "weight": weights, "output": output_tensor} else: - output_tensor = reference(input_tensor) + output_tensor = reference(input_tensor, eps=eps) return {"input": input_tensor, "output": output_tensor}