use crate::simd::get_simd_backend;
#[inline]
pub fn rms_norm(x: &[f32], weight: &[f32], eps: f32, out: &mut [f32]) {
get_simd_backend().rms_norm(x, weight, eps, out);
}
#[inline]
pub fn rms_norm_inplace(x: &mut [f32], weight: &[f32], eps: f32) {
get_simd_backend().rms_norm_inplace(x, weight, eps);
}
pub fn compute_rms(x: &[f32], eps: f32) -> f32 {
let n = x.len();
let sum_sq: f32 = x.iter().map(|&v| v * v).sum();
(sum_sq / n as f32 + eps).sqrt()
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_rms_norm() {
let x = vec![1.0, 2.0, 3.0, 4.0];
let weight = vec![1.0, 1.0, 1.0, 1.0];
let mut out = vec![0.0; 4];
rms_norm(&x, &weight, 1e-6, &mut out);
assert!(out.iter().all(|&v| v.is_finite()));
assert!(out.iter().any(|&v| v != 0.0));
}
}