1use rayon::prelude::*;
2
3use crate::tensor::Tensor;
4
5pub fn matmul_f32(a: &Tensor, b_t: &Tensor) -> Tensor {
19 let m = a.rows();
20 let k = a.cols();
21 let n = b_t.rows();
22 assert_eq!(
23 b_t.cols(),
24 k,
25 "matmul shape mismatch: a is [{m},{k}], b_t is [{},{}]",
26 b_t.rows(),
27 b_t.cols()
28 );
29
30 let mut out = vec![0f32; m * n];
31 if m == 1 {
32 let a_row = a.row(0);
33 out.par_iter_mut().enumerate().for_each(|(col, out_val)| {
34 let b_row = b_t.row(col);
35 let mut acc = 0f32;
36 for i in 0..k {
37 acc += a_row[i] * b_row[i];
38 }
39 *out_val = acc;
40 });
41 } else {
42 out.par_chunks_mut(n)
43 .enumerate()
44 .for_each(|(row, out_row)| {
45 let a_row = a.row(row);
46 for (col, out_val) in out_row.iter_mut().enumerate() {
47 let b_row = b_t.row(col);
48 let mut acc = 0f32;
49 for i in 0..k {
50 acc += a_row[i] * b_row[i];
51 }
52 *out_val = acc;
53 }
54 });
55 }
56
57 Tensor::new(out, vec![m, n])
58}
59
60pub fn rms_norm(x: &[f32], weight: &[f32], eps: f32) -> Vec<f32> {
63 assert_eq!(x.len(), weight.len());
64 let mean_sq = sum_sq(x) / x.len() as f32;
65 let scale = 1.0 / (mean_sq + eps).sqrt();
66 let mut out = vec![0f32; x.len()];
67 mul3_scale(x, weight, scale, &mut out);
68 out
69}
70
71pub fn rms_norm_per_head(x: &[f32], weight: &[f32], head_dim: usize, eps: f32) -> Vec<f32> {
75 assert_eq!(weight.len(), head_dim);
76 assert_eq!(x.len() % head_dim, 0);
77 let mut out = vec![0f32; x.len()];
78 for (head, out_h) in x.chunks_exact(head_dim).zip(out.chunks_exact_mut(head_dim)) {
79 let mean_sq = sum_sq(head) / head_dim as f32;
80 let scale = 1.0 / (mean_sq + eps).sqrt();
81 mul3_scale(head, weight, scale, out_h);
82 }
83 out
84}
85
86#[inline]
87fn sum_sq(x: &[f32]) -> f32 {
88 #[cfg(target_arch = "aarch64")]
89 {
90 if std::arch::is_aarch64_feature_detected!("neon") {
91 return unsafe { sum_sq_neon(x) };
92 }
93 }
94 #[cfg(target_arch = "x86_64")]
95 {
96 if std::is_x86_feature_detected!("avx2") && std::is_x86_feature_detected!("fma") {
97 return unsafe { sum_sq_avx2(x) };
98 }
99 }
100 x.iter().map(|v| v * v).sum()
101}
102
103#[inline]
104fn mul3_scale(x: &[f32], w: &[f32], scale: f32, out: &mut [f32]) {
105 debug_assert_eq!(x.len(), w.len());
106 debug_assert_eq!(x.len(), out.len());
107 #[cfg(target_arch = "aarch64")]
108 {
109 if std::arch::is_aarch64_feature_detected!("neon") {
110 unsafe { mul3_scale_neon(x, w, scale, out) };
111 return;
112 }
113 }
114 #[cfg(target_arch = "x86_64")]
115 {
116 if std::is_x86_feature_detected!("avx2") && std::is_x86_feature_detected!("fma") {
117 unsafe { mul3_scale_avx2(x, w, scale, out) };
118 return;
119 }
120 }
121 for ((o, &xv), &wv) in out.iter_mut().zip(x).zip(w) {
122 *o = xv * scale * wv;
123 }
124}
125
126#[cfg(target_arch = "aarch64")]
127#[target_feature(enable = "neon")]
128unsafe fn sum_sq_neon(x: &[f32]) -> f32 {
129 use std::arch::aarch64::*;
130 let n = x.len();
131 let mut acc = vdupq_n_f32(0.0);
132 let mut i = 0;
133 while i + 4 <= n {
134 let v = vld1q_f32(x.as_ptr().add(i));
135 acc = vfmaq_f32(acc, v, v);
136 i += 4;
137 }
138 let mut sum = vaddvq_f32(acc);
139 while i < n {
140 sum += x[i] * x[i];
141 i += 1;
142 }
143 sum
144}
145
146#[cfg(target_arch = "aarch64")]
147#[target_feature(enable = "neon")]
148unsafe fn mul3_scale_neon(x: &[f32], w: &[f32], scale: f32, out: &mut [f32]) {
149 use std::arch::aarch64::*;
150 let n = x.len();
151 let vs = vdupq_n_f32(scale);
152 let mut i = 0;
153 while i + 4 <= n {
154 let xv = vld1q_f32(x.as_ptr().add(i));
155 let wv = vld1q_f32(w.as_ptr().add(i));
156 vst1q_f32(out.as_mut_ptr().add(i), vmulq_f32(vmulq_f32(xv, vs), wv));
157 i += 4;
158 }
159 while i < n {
160 out[i] = x[i] * scale * w[i];
161 i += 1;
162 }
163}
164
165#[cfg(target_arch = "x86_64")]
166#[target_feature(enable = "avx2,fma")]
167unsafe fn sum_sq_avx2(x: &[f32]) -> f32 {
168 use std::arch::x86_64::*;
169 let n = x.len();
170 let mut acc = _mm256_setzero_ps();
171 let mut i = 0;
172 while i + 8 <= n {
173 let v = _mm256_loadu_ps(x.as_ptr().add(i));
174 acc = _mm256_fmadd_ps(v, v, acc);
175 i += 8;
176 }
177 let lo = _mm256_castps256_ps128(acc);
178 let hi = _mm256_extractf128_ps(acc, 1);
179 let mut s128 = _mm_add_ps(lo, hi);
180 s128 = _mm_add_ps(s128, _mm_movehl_ps(s128, s128));
181 s128 = _mm_add_ss(s128, _mm_shuffle_ps(s128, s128, 1));
182 let mut sum = _mm_cvtss_f32(s128);
183 while i < n {
184 sum += x[i] * x[i];
185 i += 1;
186 }
187 sum
188}
189
190#[cfg(target_arch = "x86_64")]
191#[target_feature(enable = "avx2,fma")]
192unsafe fn mul3_scale_avx2(x: &[f32], w: &[f32], scale: f32, out: &mut [f32]) {
193 use std::arch::x86_64::*;
194 let n = x.len();
195 let vs = _mm256_set1_ps(scale);
196 let mut i = 0;
197 while i + 8 <= n {
198 let xv = _mm256_loadu_ps(x.as_ptr().add(i));
199 let wv = _mm256_loadu_ps(w.as_ptr().add(i));
200 _mm256_storeu_ps(
201 out.as_mut_ptr().add(i),
202 _mm256_mul_ps(_mm256_mul_ps(xv, vs), wv),
203 );
204 i += 8;
205 }
206 while i < n {
207 out[i] = x[i] * scale * w[i];
208 i += 1;
209 }
210}
211
212pub fn softcap_inplace(x: &mut [f32], softcap: f32) {
215 if softcap <= 0.0 {
216 return;
217 }
218 let inv = 1.0 / softcap;
219 for v in x.iter_mut() {
220 *v = softcap * (*v * inv).tanh();
221 }
222}
223
224pub fn gelu(x: f32) -> f32 {
226 const K: f32 = 0.797_884_6; const C: f32 = 0.044_715;
229 0.5 * x * (1.0 + (K * (x + C * x * x * x)).tanh())
230}
231
232pub fn geglu(gate: &[f32], up: &[f32]) -> Vec<f32> {
252 assert_eq!(gate.len(), up.len());
253 par_gated_chunks(gate, up, gelu_mul)
254}
255
256const GATED_PAR_MIN: usize = 1 << 15;
261
262#[inline]
277fn par_gated_chunks<F>(gate: &[f32], up: &[f32], f: F) -> Vec<f32>
278where
279 F: Fn(&[f32], &[f32], &mut [f32]) + Sync + Send,
280{
281 let n = gate.len();
282 let mut out = vec![0f32; n];
283 if n < GATED_PAR_MIN {
284 f(gate, up, &mut out);
285 return out;
286 }
287 let chunk = (n.div_ceil(rayon::current_num_threads() * 4)).next_multiple_of(16);
289 out.par_chunks_mut(chunk)
290 .zip(gate.par_chunks(chunk))
291 .zip(up.par_chunks(chunk))
292 .for_each(|((o, g), u)| f(g, u, o));
293 out
294}
295
296fn gelu_mul(gate: &[f32], up: &[f32], out: &mut [f32]) {
299 debug_assert_eq!(gate.len(), up.len());
300 debug_assert_eq!(gate.len(), out.len());
301 #[cfg(target_arch = "aarch64")]
302 {
303 if std::arch::is_aarch64_feature_detected!("neon") {
304 unsafe { gelu_mul_neon(gate, up, out) };
305 return;
306 }
307 }
308 #[cfg(target_arch = "x86_64")]
309 {
310 if std::is_x86_feature_detected!("avx2") && std::is_x86_feature_detected!("fma") {
311 unsafe { gelu_mul_avx2(gate, up, out) };
312 return;
313 }
314 }
315 for ((o, g), u) in out.iter_mut().zip(gate.iter()).zip(up.iter()) {
316 *o = gelu(*g) * *u;
317 }
318}
319
320fn silu_mul(gate: &[f32], up: &[f32], out: &mut [f32]) {
322 debug_assert_eq!(gate.len(), up.len());
323 debug_assert_eq!(gate.len(), out.len());
324 #[cfg(target_arch = "aarch64")]
325 {
326 if std::arch::is_aarch64_feature_detected!("neon") {
327 unsafe { silu_mul_neon(gate, up, out) };
328 return;
329 }
330 }
331 #[cfg(target_arch = "x86_64")]
332 {
333 if std::is_x86_feature_detected!("avx2") && std::is_x86_feature_detected!("fma") {
334 unsafe { silu_mul_avx2(gate, up, out) };
335 return;
336 }
337 }
338 for ((o, g), u) in out.iter_mut().zip(gate.iter()).zip(up.iter()) {
339 *o = silu(*g) * *u;
340 }
341}
342
343#[cfg(any(target_arch = "aarch64", target_arch = "x86_64"))]
354mod exp_consts {
355 pub use crate::vexp::*;
360
361 pub const EXP_CLAMP: f32 = 87.0;
387 pub const GELU_K: f32 = 0.797_884_6;
389 pub const GELU_A: f32 = -2.0 * GELU_K;
393 pub const GELU_B: f32 = -2.0 * GELU_K * 0.044_715;
395}
396#[cfg(any(target_arch = "aarch64", target_arch = "x86_64"))]
397use exp_consts::*;
398
399#[cfg(target_arch = "aarch64")]
412#[target_feature(enable = "neon")]
413#[inline]
414unsafe fn expf_neon(x: std::arch::aarch64::float32x4_t) -> std::arch::aarch64::float32x4_t {
415 use std::arch::aarch64::*;
416 let x = vminq_f32(
417 vmaxq_f32(x, vdupq_n_f32(-EXP_CLAMP)),
418 vdupq_n_f32(EXP_CLAMP),
419 );
420 let r = vdupq_n_f32(EXP_SHIFT);
421 let z = vfmaq_f32(r, x, vdupq_n_f32(EXP_LOG2E));
422 let n = vsubq_f32(z, r);
423 let b = vfmsq_f32(
425 vfmsq_f32(x, n, vdupq_n_f32(EXP_LN2_HI)),
426 n,
427 vdupq_n_f32(EXP_LN2_LO),
428 );
429 let e = vshlq_n_u32::<23>(vreinterpretq_u32_f32(z));
430 let k = vreinterpretq_f32_u32(vaddq_u32(e, vreinterpretq_u32_f32(vdupq_n_f32(1.0))));
431 let u = vmulq_f32(b, b);
432 let j = vfmaq_f32(
433 vmulq_f32(vdupq_n_f32(EXP_C0), b),
434 vfmaq_f32(
435 vfmaq_f32(vdupq_n_f32(EXP_C1), vdupq_n_f32(EXP_C2), b),
436 vfmaq_f32(vdupq_n_f32(EXP_C3), vdupq_n_f32(EXP_C4), b),
437 u,
438 ),
439 u,
440 );
441 vfmaq_f32(k, j, k)
442}
443
444#[cfg(any(target_arch = "aarch64", target_arch = "x86_64"))]
454#[inline]
455fn expf_scalar(x: f32) -> f32 {
456 let x = x.clamp(-EXP_CLAMP, EXP_CLAMP);
457 let z = x.mul_add(EXP_LOG2E, EXP_SHIFT);
458 let n = z - EXP_SHIFT;
459 let b = (-n).mul_add(EXP_LN2_LO, (-n).mul_add(EXP_LN2_HI, x));
460 let k = f32::from_bits((z.to_bits() << 23).wrapping_add(1.0f32.to_bits()));
461 let u = b * b;
462 let j = EXP_C4
463 .mul_add(b, EXP_C3)
464 .mul_add(u, EXP_C2.mul_add(b, EXP_C1))
465 .mul_add(u, EXP_C0 * b);
466 j.mul_add(k, k)
467}
468
469#[cfg(target_arch = "x86_64")]
472#[target_feature(enable = "avx2,fma")]
473#[inline]
474unsafe fn expf_avx2(x: std::arch::x86_64::__m256) -> std::arch::x86_64::__m256 {
475 use std::arch::x86_64::*;
476 let x = _mm256_min_ps(
477 _mm256_max_ps(x, _mm256_set1_ps(-EXP_CLAMP)),
478 _mm256_set1_ps(EXP_CLAMP),
479 );
480 let r = _mm256_set1_ps(EXP_SHIFT);
481 let z = _mm256_fmadd_ps(x, _mm256_set1_ps(EXP_LOG2E), r);
482 let n = _mm256_sub_ps(z, r);
483 let b = _mm256_fnmadd_ps(
484 n,
485 _mm256_set1_ps(EXP_LN2_LO),
486 _mm256_fnmadd_ps(n, _mm256_set1_ps(EXP_LN2_HI), x),
487 );
488 let e = _mm256_slli_epi32::<23>(_mm256_castps_si256(z));
489 let k = _mm256_castsi256_ps(_mm256_add_epi32(
490 e,
491 _mm256_castps_si256(_mm256_set1_ps(1.0)),
492 ));
493 let u = _mm256_mul_ps(b, b);
494 let j = _mm256_fmadd_ps(
495 _mm256_fmadd_ps(
496 _mm256_fmadd_ps(_mm256_set1_ps(EXP_C4), b, _mm256_set1_ps(EXP_C3)),
497 u,
498 _mm256_fmadd_ps(_mm256_set1_ps(EXP_C2), b, _mm256_set1_ps(EXP_C1)),
499 ),
500 u,
501 _mm256_mul_ps(_mm256_set1_ps(EXP_C0), b),
502 );
503 _mm256_fmadd_ps(j, k, k)
504}
505
506#[cfg(any(target_arch = "aarch64", target_arch = "x86_64"))]
511#[inline]
512fn gate_by_exp_scalar(x: f32, t: f32) -> f32 {
513 if t >= EXP_CLAMP {
514 0.0
515 } else {
516 x / (1.0 + expf_scalar(t))
517 }
518}
519
520#[cfg(any(target_arch = "aarch64", target_arch = "x86_64"))]
523#[inline]
524fn gelu_exp_arg(g: f32) -> f32 {
525 g * GELU_B.mul_add(g * g, GELU_A)
526}
527
528#[cfg(target_arch = "aarch64")]
529#[target_feature(enable = "neon")]
530unsafe fn gelu_mul_neon(gate: &[f32], up: &[f32], out: &mut [f32]) {
531 use std::arch::aarch64::*;
532 let n = out.len();
533 let nv = n & !3;
534 let one = vdupq_n_f32(1.0);
535 let zero = vdupq_n_f32(0.0);
536 let a = vdupq_n_f32(GELU_A);
537 let b = vdupq_n_f32(GELU_B);
538 let sat = vdupq_n_f32(EXP_CLAMP);
539 let mut i = 0;
540 while i < nv {
541 let g = vld1q_f32(gate.as_ptr().add(i));
542 let t = vmulq_f32(g, vfmaq_f32(a, b, vmulq_f32(g, g)));
544 let y = vdivq_f32(g, vaddq_f32(one, expf_neon(t)));
545 let y = vbslq_f32(vcgeq_f32(t, sat), zero, y);
546 vst1q_f32(
547 out.as_mut_ptr().add(i),
548 vmulq_f32(y, vld1q_f32(up.as_ptr().add(i))),
549 );
550 i += 4;
551 }
552 for j in nv..n {
553 let g = *gate.get_unchecked(j);
554 *out.get_unchecked_mut(j) = gate_by_exp_scalar(g, gelu_exp_arg(g)) * *up.get_unchecked(j);
555 }
556}
557
558#[cfg(target_arch = "aarch64")]
559#[target_feature(enable = "neon")]
560unsafe fn silu_mul_neon(gate: &[f32], up: &[f32], out: &mut [f32]) {
561 use std::arch::aarch64::*;
562 let n = out.len();
563 let nv = n & !3;
564 let one = vdupq_n_f32(1.0);
565 let zero = vdupq_n_f32(0.0);
566 let sat = vdupq_n_f32(EXP_CLAMP);
567 let mut i = 0;
568 while i < nv {
569 let g = vld1q_f32(gate.as_ptr().add(i));
570 let t = vnegq_f32(g);
571 let y = vdivq_f32(g, vaddq_f32(one, expf_neon(t)));
572 let y = vbslq_f32(vcgeq_f32(t, sat), zero, y);
573 vst1q_f32(
574 out.as_mut_ptr().add(i),
575 vmulq_f32(y, vld1q_f32(up.as_ptr().add(i))),
576 );
577 i += 4;
578 }
579 for j in nv..n {
580 let g = *gate.get_unchecked(j);
581 *out.get_unchecked_mut(j) = gate_by_exp_scalar(g, -g) * *up.get_unchecked(j);
582 }
583}
584
585#[cfg(target_arch = "x86_64")]
586#[target_feature(enable = "avx2,fma")]
587unsafe fn gelu_mul_avx2(gate: &[f32], up: &[f32], out: &mut [f32]) {
588 use std::arch::x86_64::*;
589 let n = out.len();
590 let nv = n & !7;
591 let one = _mm256_set1_ps(1.0);
592 let zero = _mm256_setzero_ps();
593 let a = _mm256_set1_ps(GELU_A);
594 let b = _mm256_set1_ps(GELU_B);
595 let sat = _mm256_set1_ps(EXP_CLAMP);
596 let mut i = 0;
597 while i < nv {
598 let g = _mm256_loadu_ps(gate.as_ptr().add(i));
599 let t = _mm256_mul_ps(g, _mm256_fmadd_ps(b, _mm256_mul_ps(g, g), a));
600 let y = _mm256_div_ps(g, _mm256_add_ps(one, expf_avx2(t)));
601 let y = _mm256_blendv_ps(y, zero, _mm256_cmp_ps::<_CMP_GE_OQ>(t, sat));
602 _mm256_storeu_ps(
603 out.as_mut_ptr().add(i),
604 _mm256_mul_ps(y, _mm256_loadu_ps(up.as_ptr().add(i))),
605 );
606 i += 8;
607 }
608 for j in nv..n {
609 let g = *gate.get_unchecked(j);
610 *out.get_unchecked_mut(j) = gate_by_exp_scalar(g, gelu_exp_arg(g)) * *up.get_unchecked(j);
611 }
612}
613
614#[cfg(target_arch = "x86_64")]
615#[target_feature(enable = "avx2,fma")]
616unsafe fn silu_mul_avx2(gate: &[f32], up: &[f32], out: &mut [f32]) {
617 use std::arch::x86_64::*;
618 let n = out.len();
619 let nv = n & !7;
620 let one = _mm256_set1_ps(1.0);
621 let zero = _mm256_setzero_ps();
622 let neg = _mm256_set1_ps(-0.0);
623 let sat = _mm256_set1_ps(EXP_CLAMP);
624 let mut i = 0;
625 while i < nv {
626 let g = _mm256_loadu_ps(gate.as_ptr().add(i));
627 let t = _mm256_xor_ps(g, neg);
630 let y = _mm256_div_ps(g, _mm256_add_ps(one, expf_avx2(t)));
631 let y = _mm256_blendv_ps(y, zero, _mm256_cmp_ps::<_CMP_GE_OQ>(t, sat));
632 _mm256_storeu_ps(
633 out.as_mut_ptr().add(i),
634 _mm256_mul_ps(y, _mm256_loadu_ps(up.as_ptr().add(i))),
635 );
636 i += 8;
637 }
638 for j in nv..n {
639 let g = *gate.get_unchecked(j);
640 *out.get_unchecked_mut(j) = gate_by_exp_scalar(g, -g) * *up.get_unchecked(j);
641 }
642}
643
644pub fn layer_norm(x: &[f32], weight: &[f32], bias: &[f32], eps: f32) -> Vec<f32> {
656 assert_eq!(x.len(), weight.len());
657 assert_eq!(x.len(), bias.len());
658 let n = x.len() as f32;
659 let mean = x.iter().sum::<f32>() / n;
660 let var = x.iter().map(|v| (v - mean).powi(2)).sum::<f32>() / n;
661 let inv_std = 1.0 / (var + eps).sqrt();
662 x.iter()
663 .zip(weight.iter())
664 .zip(bias.iter())
665 .map(|((v, w), b)| (v - mean) * inv_std * w + b)
666 .collect()
667}
668
669pub fn silu(x: f32) -> f32 {
672 x / (1.0 + (-x).exp())
673}
674
675pub fn swiglu(gate: &[f32], up: &[f32]) -> Vec<f32> {
685 assert_eq!(gate.len(), up.len());
686 par_gated_chunks(gate, up, silu_mul)
687}
688
689pub fn situ_and_mul(gate: &[f32], up: &[f32], beta: f32, linear_beta: f32) -> Vec<f32> {
698 assert_eq!(gate.len(), up.len());
699 gate.iter()
700 .zip(up.iter())
701 .map(|(g, u)| {
702 let situ_a = beta * (g / beta).tanh() * (1.0 / (1.0 + (-g).exp()));
703 let up_t = linear_beta * (u / linear_beta).tanh();
704 situ_a * up_t
705 })
706 .collect()
707}
708
709#[cfg(test)]
710mod tests {
711 use super::*;
712
713 #[test]
727 fn parallel_gated_activations_are_bit_identical_to_the_serial_form() {
728 for n in [7usize, GATED_PAR_MIN - 1, GATED_PAR_MIN, 300_007] {
734 let gate: Vec<f32> = (0..n)
735 .map(|i| ((i as f32) * 0.0037 - 4.0).sin() * 6.0)
736 .collect();
737 let up: Vec<f32> = (0..n)
738 .map(|i| ((i as f32) * 0.0041 + 1.0).cos() * 2.5)
739 .collect();
740
741 let one_at_a_time = |f: fn(&[f32], &[f32], &mut [f32])| -> Vec<f32> {
742 let mut out = vec![0f32; n];
743 for i in 0..n {
744 f(&gate[i..i + 1], &up[i..i + 1], &mut out[i..i + 1]);
745 }
746 out
747 };
748 assert_eq!(
749 swiglu(&gate, &up),
750 one_at_a_time(silu_mul),
751 "swiglu at n = {n}"
752 );
753 assert_eq!(
754 geglu(&gate, &up),
755 one_at_a_time(gelu_mul),
756 "geglu at n = {n}"
757 );
758 }
759 }
760
761 #[test]
795 fn geglu_and_swiglu_are_no_less_accurate_than_the_libm_forms_they_replace() {
796 fn sweep(
804 reference: fn(f64) -> f64,
805 vector: fn(f32) -> f32,
806 libm: fn(f32) -> f32,
807 ) -> (f64, f64, u32, u32) {
808 let (mut v_abs, mut l_abs) = (0f64, 0f64);
809 let (mut v_lost, mut l_lost) = (0u32, 0u32);
810 let mut i = -120_000i32;
811 while i <= 120_000 {
812 let x = i as f32 * 0.001;
813 let want = reference(x as f64);
814 let (v, l) = (vector(x) as f64 - want, libm(x) as f64 - want);
815 let scale = (x as f64).abs().max(1e-30);
816 v_abs = v_abs.max(v.abs() / scale);
817 l_abs = l_abs.max(l.abs() / scale);
818 if want.abs() >= 1e-30 {
819 v_lost += u32::from(v.abs() / want.abs() > 1e-3);
820 l_lost += u32::from(l.abs() / want.abs() > 1e-3);
821 }
822 i += 1;
823 }
824 (v_abs, l_abs, v_lost, l_lost)
825 }
826
827 fn gelu_f64(x: f64) -> f64 {
831 const K: f64 = 0.797_884_560_802_865_4;
832 let u = K * (x + 0.044_715 * x * x * x);
833 x / (1.0 + (-2.0 * u).exp())
834 }
835 fn silu_f64(x: f64) -> f64 {
836 x / (1.0 + (-x).exp())
837 }
838 fn vec_gelu(x: f32) -> f32 {
839 let mut out = [0f32; 1];
840 gelu_mul(&[x], &[1.0], &mut out);
841 out[0]
842 }
843 fn vec_silu(x: f32) -> f32 {
844 let mut out = [0f32; 1];
845 silu_mul(&[x], &[1.0], &mut out);
846 out[0]
847 }
848
849 for (what, reference, vector, libm) in [
850 (
851 "GELU",
852 gelu_f64 as fn(f64) -> f64,
853 vec_gelu as fn(f32) -> f32,
854 gelu as fn(f32) -> f32,
855 ),
856 ("SiLU", silu_f64, vec_silu, silu),
857 ] {
858 let (v_abs, l_abs, v_lost, l_lost) = sweep(reference, vector, libm);
859 assert_eq!(
860 v_lost, 0,
861 "vector {what} lost more than 0.1% of the value at {v_lost} samples \
862 (the form it replaces: {l_lost})"
863 );
864 assert!(
865 v_lost <= l_lost,
866 "vector {what} loses values the form it replaces kept: {v_lost} vs {l_lost}"
867 );
868 assert!(
869 v_abs <= l_abs * 1.25,
870 "vector {what} contributes more absolute error than the form it \
871 replaces: {v_abs:e} vs {l_abs:e}"
872 );
873 }
874 }
875
876 #[test]
887 fn the_two_sided_clamp_leaves_the_saturating_tails_correct() {
888 fn gelu_f64(x: f64) -> f64 {
889 const K: f64 = 0.797_884_560_802_865_4;
890 let u = K * (x + 0.044_715 * x * x * x);
891 x / (1.0 + (-2.0 * u).exp())
892 }
893 for x in [
894 -1e30f32, -1e10, -1000.0, -120.0, -88.0, -12.0, -10.0, 10.0, 12.0, 20.0, 120.0, 1000.0,
895 1e30,
896 ] {
897 let mut g = [0f32; 1];
898 gelu_mul(&[x], &[1.0], &mut g);
899 let mut s = [0f32; 1];
900 silu_mul(&[x], &[1.0], &mut s);
901 assert!(g[0].is_finite() || x.abs() > 1e20, "gelu({x}) = {}", g[0]);
902 assert!(s[0].is_finite() || x.abs() > 1e20, "silu({x}) = {}", s[0]);
903
904 let want_g = gelu_f64(x as f64);
907 let want_s = (x as f64) / (1.0 + (-(x as f64)).exp());
908 assert!(
909 (g[0] as f64 - want_g).abs() <= 1e-30 + 1e-6 * want_g.abs(),
910 "gelu({x}) = {} want {want_g:e}",
911 g[0]
912 );
913 assert!(
914 (s[0] as f64 - want_s).abs() <= 1e-30 + 1e-6 * want_s.abs(),
915 "silu({x}) = {} want {want_s:e}",
916 s[0]
917 );
918 if x >= 10.0 {
923 assert_eq!(g[0], x, "gelu should saturate to the identity at {x}");
924 }
925 if x >= 20.0 {
926 assert_eq!(s[0], x, "silu should saturate to the identity at {x}");
927 }
928 }
929 }
930
931 #[test]
932 fn layer_norm_zero_mean_unit_var_input_is_unchanged_by_weight_one_bias_zero() {
933 let x = vec![-1.0, 1.0];
936 let weight = vec![1.0, 1.0];
937 let bias = vec![0.0, 0.0];
938 let out = layer_norm(&x, &weight, &bias, 0.0);
939 assert!((out[0] - (-1.0)).abs() < 1e-4);
940 assert!((out[1] - 1.0).abs() < 1e-4);
941 }
942
943 #[test]
944 fn layer_norm_applies_affine_weight_and_bias_after_normalizing() {
945 let x = vec![-1.0, 1.0];
946 let weight = vec![2.0, 3.0];
947 let bias = vec![10.0, -10.0];
948 let out = layer_norm(&x, &weight, &bias, 0.0);
949 assert!((out[0] - (-2.0 + 10.0)).abs() < 1e-4);
951 assert!((out[1] - (3.0 - 10.0)).abs() < 1e-4);
952 }
953
954 #[test]
955 fn layer_norm_constant_input_is_zero_before_bias() {
956 let x = vec![5.0, 5.0, 5.0];
959 let weight = vec![1.0, 1.0, 1.0];
960 let bias = vec![0.25, 0.25, 0.25];
961 let out = layer_norm(&x, &weight, &bias, 1e-5);
962 for v in out {
963 assert!((v - 0.25).abs() < 1e-4);
964 }
965 }
966
967 #[test]
968 fn matmul_identity_returns_input() {
969 let a = Tensor::new(vec![1.0, 2.0, 3.0, 4.0], vec![2, 2]);
971 let identity = Tensor::new(vec![1.0, 0.0, 0.0, 1.0], vec![2, 2]);
972 let out = matmul_f32(&a, &identity);
973 assert_eq!(out.data, vec![1.0, 2.0, 3.0, 4.0]);
974 }
975
976 #[test]
977 fn matmul_known_values() {
978 let a = Tensor::new(vec![1.0, 2.0, 3.0], vec![1, 3]);
980 let b_t = Tensor::new(vec![1.0, 1.0, 1.0], vec![1, 3]);
981 let out = matmul_f32(&a, &b_t);
982 assert_eq!(out.shape, vec![1, 1]);
983 assert_eq!(out.data[0], 6.0);
984 }
985
986 #[test]
987 fn matmul_single_row_batch_matches_sequential_dot_products() {
988 let a = Tensor::new(vec![1.0, 2.0, 3.0], vec![1, 3]);
992 let b_t = Tensor::new(
993 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],
994 vec![4, 3],
995 );
996 let out = matmul_f32(&a, &b_t);
997 assert_eq!(out.shape, vec![1, 4]);
998 assert_eq!(out.data, vec![1.0, 2.0, 6.0, 2.0]);
999 }
1000
1001 #[test]
1002 fn rms_norm_unit_weight_preserves_direction() {
1003 let x = vec![3.0, 4.0];
1004 let w = vec![1.0, 1.0];
1005 let out = rms_norm(&x, &w, 1e-6);
1006 assert!((out[0] / out[1] - 3.0 / 4.0).abs() < 1e-4);
1008 }
1009
1010 #[test]
1011 fn silu_is_zero_at_zero_and_monotonic_ish() {
1012 assert!((silu(0.0)).abs() < 1e-6);
1013 assert!(silu(5.0) > silu(1.0));
1014 }
1015
1016 #[test]
1022 fn situ_and_mul_matches_independent_python_reference() {
1023 let cases = [
1024 (0.0f32, 0.0f32, 0.0f32),
1025 (2.0, -3.0, -4.861_066_3),
1026 (-1.5, 10.0, -2.483_860_7),
1027 ];
1028 for (gate, up, expected) in cases {
1029 let got = situ_and_mul(&[gate], &[up], 4.0, 25.0)[0];
1030 assert!(
1031 (got - expected).abs() < 1e-4,
1032 "situ({gate},{up}): rust={got} python={expected}"
1033 );
1034 }
1035 }
1036}