1use half::bf16;
4
5pub 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
17pub 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 if v > 20.0 { v } else { v.exp().ln_1p() }
29}
30
31pub 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
48pub 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}