use rayon::prelude::*;
use crate::tensor::Tensor;
pub fn matmul_f32(a: &Tensor, b_t: &Tensor) -> Tensor {
let m = a.rows();
let k = a.cols();
let n = b_t.rows();
assert_eq!(
b_t.cols(),
k,
"matmul shape mismatch: a is [{m},{k}], b_t is [{},{}]",
b_t.rows(),
b_t.cols()
);
let mut out = vec![0f32; m * n];
if m == 1 {
let a_row = a.row(0);
out.par_iter_mut().enumerate().for_each(|(col, out_val)| {
let b_row = b_t.row(col);
let mut acc = 0f32;
for i in 0..k {
acc += a_row[i] * b_row[i];
}
*out_val = acc;
});
} else {
out.par_chunks_mut(n)
.enumerate()
.for_each(|(row, out_row)| {
let a_row = a.row(row);
for (col, out_val) in out_row.iter_mut().enumerate() {
let b_row = b_t.row(col);
let mut acc = 0f32;
for i in 0..k {
acc += a_row[i] * b_row[i];
}
*out_val = acc;
}
});
}
Tensor::new(out, vec![m, n])
}
pub fn rms_norm(x: &[f32], weight: &[f32], eps: f32) -> Vec<f32> {
assert_eq!(x.len(), weight.len());
let mean_sq = sum_sq(x) / x.len() as f32;
let scale = 1.0 / (mean_sq + eps).sqrt();
let mut out = vec![0f32; x.len()];
mul3_scale(x, weight, scale, &mut out);
out
}
pub fn rms_norm_per_head(x: &[f32], weight: &[f32], head_dim: usize, eps: f32) -> Vec<f32> {
assert_eq!(weight.len(), head_dim);
assert_eq!(x.len() % head_dim, 0);
let mut out = vec![0f32; x.len()];
for (head, out_h) in x.chunks_exact(head_dim).zip(out.chunks_exact_mut(head_dim)) {
let mean_sq = sum_sq(head) / head_dim as f32;
let scale = 1.0 / (mean_sq + eps).sqrt();
mul3_scale(head, weight, scale, out_h);
}
out
}
#[inline]
fn sum_sq(x: &[f32]) -> f32 {
#[cfg(target_arch = "aarch64")]
{
if std::arch::is_aarch64_feature_detected!("neon") {
return unsafe { sum_sq_neon(x) };
}
}
#[cfg(target_arch = "x86_64")]
{
if std::is_x86_feature_detected!("avx2") && std::is_x86_feature_detected!("fma") {
return unsafe { sum_sq_avx2(x) };
}
}
x.iter().map(|v| v * v).sum()
}
#[inline]
fn mul3_scale(x: &[f32], w: &[f32], scale: f32, out: &mut [f32]) {
debug_assert_eq!(x.len(), w.len());
debug_assert_eq!(x.len(), out.len());
#[cfg(target_arch = "aarch64")]
{
if std::arch::is_aarch64_feature_detected!("neon") {
unsafe { mul3_scale_neon(x, w, scale, out) };
return;
}
}
#[cfg(target_arch = "x86_64")]
{
if std::is_x86_feature_detected!("avx2") && std::is_x86_feature_detected!("fma") {
unsafe { mul3_scale_avx2(x, w, scale, out) };
return;
}
}
for ((o, &xv), &wv) in out.iter_mut().zip(x).zip(w) {
*o = xv * scale * wv;
}
}
#[cfg(target_arch = "aarch64")]
#[target_feature(enable = "neon")]
unsafe fn sum_sq_neon(x: &[f32]) -> f32 {
use std::arch::aarch64::*;
let n = x.len();
let mut acc = vdupq_n_f32(0.0);
let mut i = 0;
while i + 4 <= n {
let v = vld1q_f32(x.as_ptr().add(i));
acc = vfmaq_f32(acc, v, v);
i += 4;
}
let mut sum = vaddvq_f32(acc);
while i < n {
sum += x[i] * x[i];
i += 1;
}
sum
}
#[cfg(target_arch = "aarch64")]
#[target_feature(enable = "neon")]
unsafe fn mul3_scale_neon(x: &[f32], w: &[f32], scale: f32, out: &mut [f32]) {
use std::arch::aarch64::*;
let n = x.len();
let vs = vdupq_n_f32(scale);
let mut i = 0;
while i + 4 <= n {
let xv = vld1q_f32(x.as_ptr().add(i));
let wv = vld1q_f32(w.as_ptr().add(i));
vst1q_f32(out.as_mut_ptr().add(i), vmulq_f32(vmulq_f32(xv, vs), wv));
i += 4;
}
while i < n {
out[i] = x[i] * scale * w[i];
i += 1;
}
}
#[cfg(target_arch = "x86_64")]
#[target_feature(enable = "avx2,fma")]
unsafe fn sum_sq_avx2(x: &[f32]) -> f32 {
use std::arch::x86_64::*;
let n = x.len();
let mut acc = _mm256_setzero_ps();
let mut i = 0;
while i + 8 <= n {
let v = _mm256_loadu_ps(x.as_ptr().add(i));
acc = _mm256_fmadd_ps(v, v, acc);
i += 8;
}
let lo = _mm256_castps256_ps128(acc);
let hi = _mm256_extractf128_ps(acc, 1);
let mut s128 = _mm_add_ps(lo, hi);
s128 = _mm_add_ps(s128, _mm_movehl_ps(s128, s128));
s128 = _mm_add_ss(s128, _mm_shuffle_ps(s128, s128, 1));
let mut sum = _mm_cvtss_f32(s128);
while i < n {
sum += x[i] * x[i];
i += 1;
}
sum
}
#[cfg(target_arch = "x86_64")]
#[target_feature(enable = "avx2,fma")]
unsafe fn mul3_scale_avx2(x: &[f32], w: &[f32], scale: f32, out: &mut [f32]) {
use std::arch::x86_64::*;
let n = x.len();
let vs = _mm256_set1_ps(scale);
let mut i = 0;
while i + 8 <= n {
let xv = _mm256_loadu_ps(x.as_ptr().add(i));
let wv = _mm256_loadu_ps(w.as_ptr().add(i));
_mm256_storeu_ps(
out.as_mut_ptr().add(i),
_mm256_mul_ps(_mm256_mul_ps(xv, vs), wv),
);
i += 8;
}
while i < n {
out[i] = x[i] * scale * w[i];
i += 1;
}
}
pub fn softcap_inplace(x: &mut [f32], softcap: f32) {
if softcap <= 0.0 {
return;
}
let inv = 1.0 / softcap;
for v in x.iter_mut() {
*v = softcap * (*v * inv).tanh();
}
}
pub fn gelu(x: f32) -> f32 {
const K: f32 = 0.797_884_6; const C: f32 = 0.044_715;
0.5 * x * (1.0 + (K * (x + C * x * x * x)).tanh())
}
pub fn geglu(gate: &[f32], up: &[f32]) -> Vec<f32> {
assert_eq!(gate.len(), up.len());
par_gated_chunks(gate, up, gelu_mul)
}
const GATED_PAR_MIN: usize = 1 << 15;
#[inline]
fn par_gated_chunks<F>(gate: &[f32], up: &[f32], f: F) -> Vec<f32>
where
F: Fn(&[f32], &[f32], &mut [f32]) + Sync + Send,
{
let n = gate.len();
let mut out = vec![0f32; n];
if n < GATED_PAR_MIN {
f(gate, up, &mut out);
return out;
}
let chunk = (n.div_ceil(rayon::current_num_threads() * 4)).next_multiple_of(16);
out.par_chunks_mut(chunk)
.zip(gate.par_chunks(chunk))
.zip(up.par_chunks(chunk))
.for_each(|((o, g), u)| f(g, u, o));
out
}
fn gelu_mul(gate: &[f32], up: &[f32], out: &mut [f32]) {
debug_assert_eq!(gate.len(), up.len());
debug_assert_eq!(gate.len(), out.len());
#[cfg(target_arch = "aarch64")]
{
if std::arch::is_aarch64_feature_detected!("neon") {
unsafe { gelu_mul_neon(gate, up, out) };
return;
}
}
#[cfg(target_arch = "x86_64")]
{
if std::is_x86_feature_detected!("avx2") && std::is_x86_feature_detected!("fma") {
unsafe { gelu_mul_avx2(gate, up, out) };
return;
}
}
for ((o, g), u) in out.iter_mut().zip(gate.iter()).zip(up.iter()) {
*o = gelu(*g) * *u;
}
}
fn silu_mul(gate: &[f32], up: &[f32], out: &mut [f32]) {
debug_assert_eq!(gate.len(), up.len());
debug_assert_eq!(gate.len(), out.len());
#[cfg(target_arch = "aarch64")]
{
if std::arch::is_aarch64_feature_detected!("neon") {
unsafe { silu_mul_neon(gate, up, out) };
return;
}
}
#[cfg(target_arch = "x86_64")]
{
if std::is_x86_feature_detected!("avx2") && std::is_x86_feature_detected!("fma") {
unsafe { silu_mul_avx2(gate, up, out) };
return;
}
}
for ((o, g), u) in out.iter_mut().zip(gate.iter()).zip(up.iter()) {
*o = silu(*g) * *u;
}
}
#[cfg(any(target_arch = "aarch64", target_arch = "x86_64"))]
mod exp_consts {
pub use crate::vexp::*;
pub const EXP_CLAMP: f32 = 87.0;
pub const GELU_K: f32 = 0.797_884_6;
pub const GELU_A: f32 = -2.0 * GELU_K;
pub const GELU_B: f32 = -2.0 * GELU_K * 0.044_715;
}
#[cfg(any(target_arch = "aarch64", target_arch = "x86_64"))]
use exp_consts::*;
#[cfg(target_arch = "aarch64")]
#[target_feature(enable = "neon")]
#[inline]
unsafe fn expf_neon(x: std::arch::aarch64::float32x4_t) -> std::arch::aarch64::float32x4_t {
use std::arch::aarch64::*;
let x = vminq_f32(
vmaxq_f32(x, vdupq_n_f32(-EXP_CLAMP)),
vdupq_n_f32(EXP_CLAMP),
);
let r = vdupq_n_f32(EXP_SHIFT);
let z = vfmaq_f32(r, x, vdupq_n_f32(EXP_LOG2E));
let n = vsubq_f32(z, r);
let b = vfmsq_f32(
vfmsq_f32(x, n, vdupq_n_f32(EXP_LN2_HI)),
n,
vdupq_n_f32(EXP_LN2_LO),
);
let e = vshlq_n_u32::<23>(vreinterpretq_u32_f32(z));
let k = vreinterpretq_f32_u32(vaddq_u32(e, vreinterpretq_u32_f32(vdupq_n_f32(1.0))));
let u = vmulq_f32(b, b);
let j = vfmaq_f32(
vmulq_f32(vdupq_n_f32(EXP_C0), b),
vfmaq_f32(
vfmaq_f32(vdupq_n_f32(EXP_C1), vdupq_n_f32(EXP_C2), b),
vfmaq_f32(vdupq_n_f32(EXP_C3), vdupq_n_f32(EXP_C4), b),
u,
),
u,
);
vfmaq_f32(k, j, k)
}
#[cfg(any(target_arch = "aarch64", target_arch = "x86_64"))]
#[inline]
fn expf_scalar(x: f32) -> f32 {
let x = x.clamp(-EXP_CLAMP, EXP_CLAMP);
let z = x.mul_add(EXP_LOG2E, EXP_SHIFT);
let n = z - EXP_SHIFT;
let b = (-n).mul_add(EXP_LN2_LO, (-n).mul_add(EXP_LN2_HI, x));
let k = f32::from_bits((z.to_bits() << 23).wrapping_add(1.0f32.to_bits()));
let u = b * b;
let j = EXP_C4
.mul_add(b, EXP_C3)
.mul_add(u, EXP_C2.mul_add(b, EXP_C1))
.mul_add(u, EXP_C0 * b);
j.mul_add(k, k)
}
#[cfg(target_arch = "x86_64")]
#[target_feature(enable = "avx2,fma")]
#[inline]
unsafe fn expf_avx2(x: std::arch::x86_64::__m256) -> std::arch::x86_64::__m256 {
use std::arch::x86_64::*;
let x = _mm256_min_ps(
_mm256_max_ps(x, _mm256_set1_ps(-EXP_CLAMP)),
_mm256_set1_ps(EXP_CLAMP),
);
let r = _mm256_set1_ps(EXP_SHIFT);
let z = _mm256_fmadd_ps(x, _mm256_set1_ps(EXP_LOG2E), r);
let n = _mm256_sub_ps(z, r);
let b = _mm256_fnmadd_ps(
n,
_mm256_set1_ps(EXP_LN2_LO),
_mm256_fnmadd_ps(n, _mm256_set1_ps(EXP_LN2_HI), x),
);
let e = _mm256_slli_epi32::<23>(_mm256_castps_si256(z));
let k = _mm256_castsi256_ps(_mm256_add_epi32(
e,
_mm256_castps_si256(_mm256_set1_ps(1.0)),
));
let u = _mm256_mul_ps(b, b);
let j = _mm256_fmadd_ps(
_mm256_fmadd_ps(
_mm256_fmadd_ps(_mm256_set1_ps(EXP_C4), b, _mm256_set1_ps(EXP_C3)),
u,
_mm256_fmadd_ps(_mm256_set1_ps(EXP_C2), b, _mm256_set1_ps(EXP_C1)),
),
u,
_mm256_mul_ps(_mm256_set1_ps(EXP_C0), b),
);
_mm256_fmadd_ps(j, k, k)
}
#[cfg(any(target_arch = "aarch64", target_arch = "x86_64"))]
#[inline]
fn gate_by_exp_scalar(x: f32, t: f32) -> f32 {
if t >= EXP_CLAMP {
0.0
} else {
x / (1.0 + expf_scalar(t))
}
}
#[cfg(any(target_arch = "aarch64", target_arch = "x86_64"))]
#[inline]
fn gelu_exp_arg(g: f32) -> f32 {
g * GELU_B.mul_add(g * g, GELU_A)
}
#[cfg(target_arch = "aarch64")]
#[target_feature(enable = "neon")]
unsafe fn gelu_mul_neon(gate: &[f32], up: &[f32], out: &mut [f32]) {
use std::arch::aarch64::*;
let n = out.len();
let nv = n & !3;
let one = vdupq_n_f32(1.0);
let zero = vdupq_n_f32(0.0);
let a = vdupq_n_f32(GELU_A);
let b = vdupq_n_f32(GELU_B);
let sat = vdupq_n_f32(EXP_CLAMP);
let mut i = 0;
while i < nv {
let g = vld1q_f32(gate.as_ptr().add(i));
let t = vmulq_f32(g, vfmaq_f32(a, b, vmulq_f32(g, g)));
let y = vdivq_f32(g, vaddq_f32(one, expf_neon(t)));
let y = vbslq_f32(vcgeq_f32(t, sat), zero, y);
vst1q_f32(
out.as_mut_ptr().add(i),
vmulq_f32(y, vld1q_f32(up.as_ptr().add(i))),
);
i += 4;
}
for j in nv..n {
let g = *gate.get_unchecked(j);
*out.get_unchecked_mut(j) = gate_by_exp_scalar(g, gelu_exp_arg(g)) * *up.get_unchecked(j);
}
}
#[cfg(target_arch = "aarch64")]
#[target_feature(enable = "neon")]
unsafe fn silu_mul_neon(gate: &[f32], up: &[f32], out: &mut [f32]) {
use std::arch::aarch64::*;
let n = out.len();
let nv = n & !3;
let one = vdupq_n_f32(1.0);
let zero = vdupq_n_f32(0.0);
let sat = vdupq_n_f32(EXP_CLAMP);
let mut i = 0;
while i < nv {
let g = vld1q_f32(gate.as_ptr().add(i));
let t = vnegq_f32(g);
let y = vdivq_f32(g, vaddq_f32(one, expf_neon(t)));
let y = vbslq_f32(vcgeq_f32(t, sat), zero, y);
vst1q_f32(
out.as_mut_ptr().add(i),
vmulq_f32(y, vld1q_f32(up.as_ptr().add(i))),
);
i += 4;
}
for j in nv..n {
let g = *gate.get_unchecked(j);
*out.get_unchecked_mut(j) = gate_by_exp_scalar(g, -g) * *up.get_unchecked(j);
}
}
#[cfg(target_arch = "x86_64")]
#[target_feature(enable = "avx2,fma")]
unsafe fn gelu_mul_avx2(gate: &[f32], up: &[f32], out: &mut [f32]) {
use std::arch::x86_64::*;
let n = out.len();
let nv = n & !7;
let one = _mm256_set1_ps(1.0);
let zero = _mm256_setzero_ps();
let a = _mm256_set1_ps(GELU_A);
let b = _mm256_set1_ps(GELU_B);
let sat = _mm256_set1_ps(EXP_CLAMP);
let mut i = 0;
while i < nv {
let g = _mm256_loadu_ps(gate.as_ptr().add(i));
let t = _mm256_mul_ps(g, _mm256_fmadd_ps(b, _mm256_mul_ps(g, g), a));
let y = _mm256_div_ps(g, _mm256_add_ps(one, expf_avx2(t)));
let y = _mm256_blendv_ps(y, zero, _mm256_cmp_ps::<_CMP_GE_OQ>(t, sat));
_mm256_storeu_ps(
out.as_mut_ptr().add(i),
_mm256_mul_ps(y, _mm256_loadu_ps(up.as_ptr().add(i))),
);
i += 8;
}
for j in nv..n {
let g = *gate.get_unchecked(j);
*out.get_unchecked_mut(j) = gate_by_exp_scalar(g, gelu_exp_arg(g)) * *up.get_unchecked(j);
}
}
#[cfg(target_arch = "x86_64")]
#[target_feature(enable = "avx2,fma")]
unsafe fn silu_mul_avx2(gate: &[f32], up: &[f32], out: &mut [f32]) {
use std::arch::x86_64::*;
let n = out.len();
let nv = n & !7;
let one = _mm256_set1_ps(1.0);
let zero = _mm256_setzero_ps();
let neg = _mm256_set1_ps(-0.0);
let sat = _mm256_set1_ps(EXP_CLAMP);
let mut i = 0;
while i < nv {
let g = _mm256_loadu_ps(gate.as_ptr().add(i));
let t = _mm256_xor_ps(g, neg);
let y = _mm256_div_ps(g, _mm256_add_ps(one, expf_avx2(t)));
let y = _mm256_blendv_ps(y, zero, _mm256_cmp_ps::<_CMP_GE_OQ>(t, sat));
_mm256_storeu_ps(
out.as_mut_ptr().add(i),
_mm256_mul_ps(y, _mm256_loadu_ps(up.as_ptr().add(i))),
);
i += 8;
}
for j in nv..n {
let g = *gate.get_unchecked(j);
*out.get_unchecked_mut(j) = gate_by_exp_scalar(g, -g) * *up.get_unchecked(j);
}
}
pub fn layer_norm(x: &[f32], weight: &[f32], bias: &[f32], eps: f32) -> Vec<f32> {
assert_eq!(x.len(), weight.len());
assert_eq!(x.len(), bias.len());
let n = x.len() as f32;
let mean = x.iter().sum::<f32>() / n;
let var = x.iter().map(|v| (v - mean).powi(2)).sum::<f32>() / n;
let inv_std = 1.0 / (var + eps).sqrt();
x.iter()
.zip(weight.iter())
.zip(bias.iter())
.map(|((v, w), b)| (v - mean) * inv_std * w + b)
.collect()
}
pub fn silu(x: f32) -> f32 {
x / (1.0 + (-x).exp())
}
pub fn swiglu(gate: &[f32], up: &[f32]) -> Vec<f32> {
assert_eq!(gate.len(), up.len());
par_gated_chunks(gate, up, silu_mul)
}
pub fn situ_and_mul(gate: &[f32], up: &[f32], beta: f32, linear_beta: f32) -> Vec<f32> {
assert_eq!(gate.len(), up.len());
gate.iter()
.zip(up.iter())
.map(|(g, u)| {
let situ_a = beta * (g / beta).tanh() * (1.0 / (1.0 + (-g).exp()));
let up_t = linear_beta * (u / linear_beta).tanh();
situ_a * up_t
})
.collect()
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn parallel_gated_activations_are_bit_identical_to_the_serial_form() {
for n in [7usize, GATED_PAR_MIN - 1, GATED_PAR_MIN, 300_007] {
let gate: Vec<f32> = (0..n)
.map(|i| ((i as f32) * 0.0037 - 4.0).sin() * 6.0)
.collect();
let up: Vec<f32> = (0..n)
.map(|i| ((i as f32) * 0.0041 + 1.0).cos() * 2.5)
.collect();
let one_at_a_time = |f: fn(&[f32], &[f32], &mut [f32])| -> Vec<f32> {
let mut out = vec![0f32; n];
for i in 0..n {
f(&gate[i..i + 1], &up[i..i + 1], &mut out[i..i + 1]);
}
out
};
assert_eq!(
swiglu(&gate, &up),
one_at_a_time(silu_mul),
"swiglu at n = {n}"
);
assert_eq!(
geglu(&gate, &up),
one_at_a_time(gelu_mul),
"geglu at n = {n}"
);
}
}
#[test]
fn geglu_and_swiglu_are_no_less_accurate_than_the_libm_forms_they_replace() {
fn sweep(
reference: fn(f64) -> f64,
vector: fn(f32) -> f32,
libm: fn(f32) -> f32,
) -> (f64, f64, u32, u32) {
let (mut v_abs, mut l_abs) = (0f64, 0f64);
let (mut v_lost, mut l_lost) = (0u32, 0u32);
let mut i = -120_000i32;
while i <= 120_000 {
let x = i as f32 * 0.001;
let want = reference(x as f64);
let (v, l) = (vector(x) as f64 - want, libm(x) as f64 - want);
let scale = (x as f64).abs().max(1e-30);
v_abs = v_abs.max(v.abs() / scale);
l_abs = l_abs.max(l.abs() / scale);
if want.abs() >= 1e-30 {
v_lost += u32::from(v.abs() / want.abs() > 1e-3);
l_lost += u32::from(l.abs() / want.abs() > 1e-3);
}
i += 1;
}
(v_abs, l_abs, v_lost, l_lost)
}
fn gelu_f64(x: f64) -> f64 {
const K: f64 = 0.797_884_560_802_865_4;
let u = K * (x + 0.044_715 * x * x * x);
x / (1.0 + (-2.0 * u).exp())
}
fn silu_f64(x: f64) -> f64 {
x / (1.0 + (-x).exp())
}
fn vec_gelu(x: f32) -> f32 {
let mut out = [0f32; 1];
gelu_mul(&[x], &[1.0], &mut out);
out[0]
}
fn vec_silu(x: f32) -> f32 {
let mut out = [0f32; 1];
silu_mul(&[x], &[1.0], &mut out);
out[0]
}
for (what, reference, vector, libm) in [
(
"GELU",
gelu_f64 as fn(f64) -> f64,
vec_gelu as fn(f32) -> f32,
gelu as fn(f32) -> f32,
),
("SiLU", silu_f64, vec_silu, silu),
] {
let (v_abs, l_abs, v_lost, l_lost) = sweep(reference, vector, libm);
assert_eq!(
v_lost, 0,
"vector {what} lost more than 0.1% of the value at {v_lost} samples \
(the form it replaces: {l_lost})"
);
assert!(
v_lost <= l_lost,
"vector {what} loses values the form it replaces kept: {v_lost} vs {l_lost}"
);
assert!(
v_abs <= l_abs * 1.25,
"vector {what} contributes more absolute error than the form it \
replaces: {v_abs:e} vs {l_abs:e}"
);
}
}
#[test]
fn the_two_sided_clamp_leaves_the_saturating_tails_correct() {
fn gelu_f64(x: f64) -> f64 {
const K: f64 = 0.797_884_560_802_865_4;
let u = K * (x + 0.044_715 * x * x * x);
x / (1.0 + (-2.0 * u).exp())
}
for x in [
-1e30f32, -1e10, -1000.0, -120.0, -88.0, -12.0, -10.0, 10.0, 12.0, 20.0, 120.0, 1000.0,
1e30,
] {
let mut g = [0f32; 1];
gelu_mul(&[x], &[1.0], &mut g);
let mut s = [0f32; 1];
silu_mul(&[x], &[1.0], &mut s);
assert!(g[0].is_finite() || x.abs() > 1e20, "gelu({x}) = {}", g[0]);
assert!(s[0].is_finite() || x.abs() > 1e20, "silu({x}) = {}", s[0]);
let want_g = gelu_f64(x as f64);
let want_s = (x as f64) / (1.0 + (-(x as f64)).exp());
assert!(
(g[0] as f64 - want_g).abs() <= 1e-30 + 1e-6 * want_g.abs(),
"gelu({x}) = {} want {want_g:e}",
g[0]
);
assert!(
(s[0] as f64 - want_s).abs() <= 1e-30 + 1e-6 * want_s.abs(),
"silu({x}) = {} want {want_s:e}",
s[0]
);
if x >= 10.0 {
assert_eq!(g[0], x, "gelu should saturate to the identity at {x}");
}
if x >= 20.0 {
assert_eq!(s[0], x, "silu should saturate to the identity at {x}");
}
}
}
#[test]
fn layer_norm_zero_mean_unit_var_input_is_unchanged_by_weight_one_bias_zero() {
let x = vec![-1.0, 1.0];
let weight = vec![1.0, 1.0];
let bias = vec![0.0, 0.0];
let out = layer_norm(&x, &weight, &bias, 0.0);
assert!((out[0] - (-1.0)).abs() < 1e-4);
assert!((out[1] - 1.0).abs() < 1e-4);
}
#[test]
fn layer_norm_applies_affine_weight_and_bias_after_normalizing() {
let x = vec![-1.0, 1.0];
let weight = vec![2.0, 3.0];
let bias = vec![10.0, -10.0];
let out = layer_norm(&x, &weight, &bias, 0.0);
assert!((out[0] - (-2.0 + 10.0)).abs() < 1e-4);
assert!((out[1] - (3.0 - 10.0)).abs() < 1e-4);
}
#[test]
fn layer_norm_constant_input_is_zero_before_bias() {
let x = vec![5.0, 5.0, 5.0];
let weight = vec![1.0, 1.0, 1.0];
let bias = vec![0.25, 0.25, 0.25];
let out = layer_norm(&x, &weight, &bias, 1e-5);
for v in out {
assert!((v - 0.25).abs() < 1e-4);
}
}
#[test]
fn matmul_identity_returns_input() {
let a = Tensor::new(vec![1.0, 2.0, 3.0, 4.0], vec![2, 2]);
let identity = Tensor::new(vec![1.0, 0.0, 0.0, 1.0], vec![2, 2]);
let out = matmul_f32(&a, &identity);
assert_eq!(out.data, vec![1.0, 2.0, 3.0, 4.0]);
}
#[test]
fn matmul_known_values() {
let a = Tensor::new(vec![1.0, 2.0, 3.0], vec![1, 3]);
let b_t = Tensor::new(vec![1.0, 1.0, 1.0], vec![1, 3]);
let out = matmul_f32(&a, &b_t);
assert_eq!(out.shape, vec![1, 1]);
assert_eq!(out.data[0], 6.0);
}
#[test]
fn matmul_single_row_batch_matches_sequential_dot_products() {
let a = Tensor::new(vec![1.0, 2.0, 3.0], vec![1, 3]);
let b_t = Tensor::new(
vec![1.0, 0.0, 0.0, 0.0, 1.0, 0.0, 1.0, 1.0, 1.0, 2.0, 0.0, 0.0],
vec![4, 3],
);
let out = matmul_f32(&a, &b_t);
assert_eq!(out.shape, vec![1, 4]);
assert_eq!(out.data, vec![1.0, 2.0, 6.0, 2.0]);
}
#[test]
fn rms_norm_unit_weight_preserves_direction() {
let x = vec![3.0, 4.0];
let w = vec![1.0, 1.0];
let out = rms_norm(&x, &w, 1e-6);
assert!((out[0] / out[1] - 3.0 / 4.0).abs() < 1e-4);
}
#[test]
fn silu_is_zero_at_zero_and_monotonic_ish() {
assert!((silu(0.0)).abs() < 1e-6);
assert!(silu(5.0) > silu(1.0));
}
#[test]
fn situ_and_mul_matches_independent_python_reference() {
let cases = [
(0.0f32, 0.0f32, 0.0f32),
(2.0, -3.0, -4.861_066_3),
(-1.5, 10.0, -2.483_860_7),
];
for (gate, up, expected) in cases {
let got = situ_and_mul(&[gate], &[up], 4.0, 25.0)[0];
assert!(
(got - expected).abs() < 1e-4,
"situ({gate},{up}): rust={got} python={expected}"
);
}
}
}