Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
34 changes: 19 additions & 15 deletions aie_kernels/aie2/rms_norm.cc
Original file line number Diff line number Diff line change
Expand Up @@ -9,10 +9,9 @@
#include <stdlib.h>

template <typename T, int N>
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<float, N> add_res = ::aie::zeros<float, N>();

int vector_chunks = cols / N;
Expand All @@ -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<T, N> inv_rms_v = ::aie::broadcast<T, N>(static_cast<T>(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<accfloat, N> inv_rms_v;
inv_rms_v.from_vector(::aie::broadcast<float, N>(inv_rms), 0);

for (int i = 0; i < vector_chunks; i++) {
::aie::vector<T, N> reg_a = ::aie::load_v<N>(input + i * N);
::aie::vector<T, N> norm_v = ::aie::mul(reg_a, inv_rms_v);
::aie::vector<T, N> out_v;
::aie::accum<accfloat, N> reg_a;
reg_a.from_vector(::aie::load_v<N>(input + i * N), 0);
reg_a = ::aie::mul(reg_a.template to_vector<float>(), inv_rms_v.template to_vector<float>());
if (input2) {
::aie::vector<T, N> reg_b = ::aie::load_v<N>(input2 + i * N);
out_v = ::aie::mul(norm_v, reg_b);
} else {
out_v = norm_v;
::aie::accum<accfloat, N> reg_b;
reg_b.from_vector(::aie::load_v<N>(input2 + i * N), 0);
reg_a = ::aie::mul(reg_a.template to_vector<float>(), reg_b.template to_vector<float>());
}
::aie::store_v(output + i * N, out_v);
::aie::store_v(output + i * N, reg_a.template to_vector<T>());
}

if (remaining > 0) {
Expand All @@ -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<bfloat16, 16>(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<bfloat16, 16>(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<bfloat16, 16>(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<bfloat16, 16>(a_in, b_in, c_out, size, epsilon);
}
}
34 changes: 19 additions & 15 deletions aie_kernels/aie2p/rms_norm.cc
Original file line number Diff line number Diff line change
Expand Up @@ -7,10 +7,9 @@
#include <stdlib.h>

template <typename T, int N>
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<float, N> add_res = ::aie::zeros<float, N>();

int vector_chunks = cols / N;
Expand All @@ -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<T, N> inv_rms_v = ::aie::broadcast<T, N>(static_cast<T>(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<accfloat, N> inv_rms_v;
inv_rms_v.from_vector(::aie::broadcast<float, N>(inv_rms), 0);

for (int i = 0; i < vector_chunks; i++) {
::aie::vector<T, N> reg_a = ::aie::load_v<N>(input + i * N);
::aie::vector<T, N> norm_v = ::aie::mul(reg_a, inv_rms_v);
::aie::vector<T, N> out_v;
::aie::accum<accfloat, N> reg_a;
reg_a.from_vector(::aie::load_v<N>(input + i * N), 0);
reg_a = ::aie::mul(reg_a.template to_vector<float>(), inv_rms_v.template to_vector<float>());
if (input2) {
::aie::vector<T, N> reg_b = ::aie::load_v<N>(input2 + i * N);
out_v = ::aie::mul(norm_v, reg_b);
} else {
out_v = norm_v;
::aie::accum<accfloat, N> reg_b;
reg_b.from_vector(::aie::load_v<N>(input2 + i * N), 0);
reg_a = ::aie::mul(reg_a.template to_vector<float>(), reg_b.template to_vector<float>());
}
::aie::store_v(output + i * N, out_v);
::aie::store_v(output + i * N, reg_a.template to_vector<T>());
}

if (remaining > 0) {
Expand All @@ -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<bfloat16, 16>(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<bfloat16, 16>(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<bfloat16, 16>(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<bfloat16, 16>(a_in, b_in, c_out, size, epsilon);
}
}
5 changes: 3 additions & 2 deletions iron/operators/rms_norm/design.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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
Expand All @@ -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)

Expand Down
5 changes: 3 additions & 2 deletions iron/operators/rms_norm/design_weighted.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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",
Expand All @@ -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)

Expand Down
6 changes: 4 additions & 2 deletions iron/operators/rms_norm/op.py
Original file line number Diff line number Diff line change
Expand Up @@ -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)

Expand Down Expand Up @@ -87,6 +88,7 @@ def get_mlir_artifact(self):
self.num_channels,
self.tile_size,
0, # trace_size
self.epsilon,
),
),
)
Expand Down Expand Up @@ -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)
17 changes: 10 additions & 7 deletions iron/operators/rms_norm/reference.py
Original file line number Diff line number Diff line change
Expand Up @@ -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}