Skip to main content

cherenkov/
nn.rs

1//! Small CPU math helpers for the f32 reference path.
2
3use half::bf16;
4
5/// RMSNorm: x_i * w_i / sqrt(mean(x^2) + eps). Computed in f32.
6pub fn rms_norm(x: &mut [f32], w: &[bf16], eps: f32) {
7    debug_assert_eq!(x.len(), w.len());
8
9    let ms = x.iter().map(|v| v * v).sum::<f32>() / x.len() as f32;
10    let inv = 1.0 / (ms + eps).sqrt();
11
12    for (v, wi) in x.iter_mut().zip(w) {
13        *v = *v * inv * wi.to_f32();
14    }
15}
16
17/// SiLU activation: x * sigmoid(x).
18pub fn silu(v: f32) -> f32 {
19    v / (1.0 + (-v).exp())
20}
21
22pub fn sigmoid(v: f32) -> f32 {
23    1.0 / (1.0 + (-v).exp())
24}
25
26pub fn softplus(v: f32) -> f32 {
27    // Numerically stable log(1 + e^v).
28    if v > 20.0 { v } else { v.exp().ln_1p() }
29}
30
31/// In-place softmax over `x`.
32pub fn softmax(x: &mut [f32]) {
33    let max = x.iter().cloned().fold(f32::MIN, f32::max);
34    let mut sum = 0.0;
35
36    for v in x.iter_mut() {
37        *v = (*v - max).exp();
38        sum += *v;
39    }
40
41    let inv = 1.0 / sum;
42
43    for v in x.iter_mut() {
44        *v *= inv;
45    }
46}
47
48/// L2-normalize with eps, in place: x / sqrt(sum(x^2) + eps).
49pub fn l2_norm(x: &mut [f32], eps: f32) {
50    let ss = x.iter().map(|v| v * v).sum::<f32>();
51    let inv = 1.0 / (ss + eps).sqrt();
52
53    for v in x.iter_mut() {
54        *v *= inv;
55    }
56}