Skip to main content

qualia_core_db/platform/compute_bridge/
reference.rs

1//! CPU reference microkernels — one per [`KernelClass`](super::kernel_class::KernelClass).
2//!
3//! These serve two roles at once (plan §3 + §5 step 4):
4//! 1. **The panel's CPU measurement** — each is the representative microkernel timed
5//!    to give the CPU's per-class throughput row.
6//! 2. **The correctness reference** — any GPU/NPU/vendor kernel for a class must
7//!    match its CPU reference within the class tolerance before it may be the
8//!    default. "A faster wrong answer is a regression."
9//!
10//! They are plain, correct, scalar/`rayon` implementations over caller-owned slices.
11//! The CPU path is always present and never hard-fails (plan §7).
12
13use rayon::prelude::*;
14
15/// `DenseLinear`: GEMV `y = W·x`, `W` row-major `n×n`. `y.len()==n`, `x.len()==n`.
16pub fn gemv(w: &[f32], x: &[f32], y: &mut [f32]) {
17    let n = x.len();
18    debug_assert_eq!(w.len(), n * n);
19    debug_assert_eq!(y.len(), n);
20    y.par_iter_mut().enumerate().for_each(|(i, o)| {
21        let row = &w[i * n..(i + 1) * n];
22        *o = row.iter().zip(x.iter()).map(|(a, b)| a * b).sum();
23    });
24}
25
26/// `ElementwiseMap`: fused `y = a·x + b` over a large vector.
27pub fn axpb(a: f32, x: &[f32], b: f32, y: &mut [f32]) {
28    debug_assert_eq!(x.len(), y.len());
29    y.par_iter_mut()
30        .zip(x.par_iter())
31        .for_each(|(o, &xi)| *o = a * xi + b);
32}
33
34/// `Reduction`: sum of a large vector (pairwise/parallel; deterministic enough for
35/// the tolerance gate).
36pub fn reduce_sum(x: &[f32]) -> f64 {
37    x.par_iter().map(|&v| v as f64).sum()
38}
39
40/// `Stencil`: 1-D 3-point Laplacian `y[i] = x[i-1] - 2·x[i] + x[i+1]`, with the
41/// ends clamped (one-sided zero). `y.len()==x.len()`.
42pub fn stencil3(x: &[f32], y: &mut [f32]) {
43    let n = x.len();
44    debug_assert_eq!(y.len(), n);
45    if n == 0 {
46        return;
47    }
48    if n == 1 {
49        y[0] = 0.0;
50        return;
51    }
52    y[0] = x[1] - x[0];
53    for i in 1..n - 1 {
54        y[i] = x[i - 1] - 2.0 * x[i] + x[i + 1];
55    }
56    y[n - 1] = x[n - 2] - x[n - 1];
57}
58
59/// `AllPairs`: total pairwise inverse-distance potential `Σ_{i<j} 1/|p_i − p_j|`
60/// over 3-D points (`pts.len()==3·n`). A representative N-body reduction.
61pub fn allpairs_potential(pts: &[f32]) -> f64 {
62    let n = pts.len() / 3;
63    (0..n)
64        .into_par_iter()
65        .map(|i| {
66            let (xi, yi, zi) = (
67                pts[3 * i] as f64,
68                pts[3 * i + 1] as f64,
69                pts[3 * i + 2] as f64,
70            );
71            let mut acc = 0.0;
72            for j in (i + 1)..n {
73                let dx = xi - pts[3 * j] as f64;
74                let dy = yi - pts[3 * j + 1] as f64;
75                let dz = zi - pts[3 * j + 2] as f64;
76                let r = (dx * dx + dy * dy + dz * dz).sqrt();
77                if r > 0.0 {
78                    acc += 1.0 / r;
79                }
80            }
81            acc
82        })
83        .sum()
84}
85
86/// `Fft`: in-place iterative radix-2 Cooley–Tukey FFT of a complex signal held as
87/// parallel `re`/`im` slices. `re.len()==im.len()` must be a power of two.
88/// `inverse=false` is the forward transform (no 1/N scaling — matches the textbook
89/// DFT the tests check against).
90pub fn fft_radix2(re: &mut [f32], im: &mut [f32], inverse: bool) {
91    let n = re.len();
92    debug_assert_eq!(im.len(), n);
93    if n <= 1 {
94        return;
95    }
96    debug_assert!(
97        n.is_power_of_two(),
98        "fft_radix2 requires a power-of-two length"
99    );
100
101    // Bit-reversal permutation.
102    let mut j = 0usize;
103    for i in 1..n {
104        let mut bit = n >> 1;
105        while j & bit != 0 {
106            j ^= bit;
107            bit >>= 1;
108        }
109        j ^= bit;
110        if i < j {
111            re.swap(i, j);
112            im.swap(i, j);
113        }
114    }
115
116    // Danielson–Lanczos butterflies.
117    let sign = if inverse { 1.0f64 } else { -1.0f64 };
118    let mut len = 2usize;
119    while len <= n {
120        let ang = sign * 2.0 * std::f64::consts::PI / len as f64;
121        let (wlen_re, wlen_im) = (ang.cos(), ang.sin());
122        let mut i = 0usize;
123        while i < n {
124            let (mut w_re, mut w_im) = (1.0f64, 0.0f64);
125            for k in 0..len / 2 {
126                let u_re = re[i + k] as f64;
127                let u_im = im[i + k] as f64;
128                let v_re = re[i + k + len / 2] as f64 * w_re - im[i + k + len / 2] as f64 * w_im;
129                let v_im = re[i + k + len / 2] as f64 * w_im + im[i + k + len / 2] as f64 * w_re;
130                re[i + k] = (u_re + v_re) as f32;
131                im[i + k] = (u_im + v_im) as f32;
132                re[i + k + len / 2] = (u_re - v_re) as f32;
133                im[i + k + len / 2] = (u_im - v_im) as f32;
134                let nw_re = w_re * wlen_re - w_im * wlen_im;
135                w_im = w_re * wlen_im + w_im * wlen_re;
136                w_re = nw_re;
137            }
138            i += len;
139        }
140        len <<= 1;
141    }
142}
143
144/// `Scan`: inclusive prefix sum `y[i] = Σ_{k≤i} x[k]`. `y.len()==x.len()`.
145pub fn prefix_sum(x: &[f32], y: &mut [f32]) {
146    debug_assert_eq!(x.len(), y.len());
147    let mut acc = 0.0f64;
148    for (o, &xi) in y.iter_mut().zip(x.iter()) {
149        acc += xi as f64;
150        *o = acc as f32;
151    }
152}
153
154/// `Divergent`: a branch-heavy Monte-Carlo step — estimate π by the fraction of
155/// `steps` deterministic-LCG samples landing inside the unit circle (×4). The branch
156/// (`inside ? ... : ...`) is the divergence this class represents. Deterministic so
157/// it is reproducible as a correctness reference.
158pub fn monte_carlo_pi(steps: usize) -> f64 {
159    let mut state = 0x2545F4914F6CDD1Du64;
160    let mut next = || {
161        state = state
162            .wrapping_mul(6364136223846793005)
163            .wrapping_add(1442695040888963407);
164        ((state >> 11) as f64) / ((1u64 << 53) as f64)
165    };
166    let mut inside = 0usize;
167    for _ in 0..steps {
168        let x = next();
169        let y = next();
170        if x * x + y * y <= 1.0 {
171            inside += 1;
172        }
173    }
174    if steps == 0 {
175        0.0
176    } else {
177        4.0 * inside as f64 / steps as f64
178    }
179}
180
181#[cfg(test)]
182mod tests {
183    use super::*;
184
185    #[test]
186    fn gemv_matches_hand() {
187        // W = [[1,2],[3,4]], x = [1,1] → y = [3,7].
188        let w = [1.0, 2.0, 3.0, 4.0];
189        let x = [1.0, 1.0];
190        let mut y = [0.0; 2];
191        gemv(&w, &x, &mut y);
192        assert!((y[0] - 3.0).abs() < 1e-5 && (y[1] - 7.0).abs() < 1e-5);
193    }
194
195    #[test]
196    fn axpb_is_affine() {
197        let x = [1.0, 2.0, 3.0];
198        let mut y = [0.0; 3];
199        axpb(2.0, &x, 1.0, &mut y);
200        assert_eq!(y, [3.0, 5.0, 7.0]);
201    }
202
203    #[test]
204    fn reduce_and_scan_agree_on_total() {
205        let x: Vec<f32> = (1..=100).map(|i| i as f32).collect();
206        let total = reduce_sum(&x);
207        let mut ps = vec![0.0f32; x.len()];
208        prefix_sum(&x, &mut ps);
209        assert!((total - 5050.0).abs() < 1e-6);
210        assert!((*ps.last().unwrap() as f64 - total).abs() < 1e-3);
211        assert!((ps[0] - 1.0).abs() < 1e-6);
212    }
213
214    #[test]
215    fn stencil_of_linear_is_zero_interior() {
216        // A linear ramp has zero second difference in the interior.
217        let x: Vec<f32> = (0..16).map(|i| 2.0 * i as f32 + 1.0).collect();
218        let mut y = vec![0.0f32; x.len()];
219        stencil3(&x, &mut y);
220        for v in &y[1..x.len() - 1] {
221            assert!(
222                v.abs() < 1e-4,
223                "interior second difference of a ramp must be ~0, got {v}"
224            );
225        }
226    }
227
228    #[test]
229    fn fft_matches_naive_dft() {
230        let n = 8usize;
231        let signal: Vec<f32> = (0..n).map(|i| (i as f32 * 0.7).sin()).collect();
232        let mut re: Vec<f32> = signal.clone();
233        let mut im = vec![0.0f32; n];
234        fft_radix2(&mut re, &mut im, false);
235        // Naive DFT reference.
236        for k in 0..n {
237            let mut dr = 0.0f64;
238            let mut di = 0.0f64;
239            for (t, &s) in signal.iter().enumerate() {
240                let ang = -2.0 * std::f64::consts::PI * (k * t) as f64 / n as f64;
241                dr += s as f64 * ang.cos();
242                di += s as f64 * ang.sin();
243            }
244            assert!(
245                (re[k] as f64 - dr).abs() < 1e-3,
246                "re[{k}] {} vs {dr}",
247                re[k]
248            );
249            assert!(
250                (im[k] as f64 - di).abs() < 1e-3,
251                "im[{k}] {} vs {di}",
252                im[k]
253            );
254        }
255    }
256
257    #[test]
258    fn fft_inverse_recovers_signal() {
259        let n = 16usize;
260        let signal: Vec<f32> = (0..n).map(|i| (i as f32).cos()).collect();
261        let mut re = signal.clone();
262        let mut im = vec![0.0f32; n];
263        fft_radix2(&mut re, &mut im, false);
264        fft_radix2(&mut re, &mut im, true);
265        for (i, &s) in signal.iter().enumerate() {
266            assert!(
267                (re[i] / n as f32 - s).abs() < 1e-3,
268                "ifft[{i}] {} vs {s}",
269                re[i] / n as f32
270            );
271        }
272    }
273
274    #[test]
275    fn monte_carlo_pi_is_in_range_and_deterministic() {
276        let a = monte_carlo_pi(1 << 16);
277        let b = monte_carlo_pi(1 << 16);
278        assert_eq!(a, b, "must be deterministic to serve as a reference");
279        assert!(
280            (a - std::f64::consts::PI).abs() < 0.1,
281            "π estimate {a} too far off"
282        );
283    }
284
285    #[test]
286    fn allpairs_two_points_is_inverse_distance() {
287        // Two points 3 apart → potential = 1/3.
288        let pts = [0.0, 0.0, 0.0, 3.0, 0.0, 0.0];
289        let pot = allpairs_potential(&pts);
290        assert!((pot - 1.0 / 3.0).abs() < 1e-6, "{pot}");
291    }
292}