qualia_core_db/platform/compute_bridge/
reference.rs1use rayon::prelude::*;
14
15pub 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
26pub 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
34pub fn reduce_sum(x: &[f32]) -> f64 {
37 x.par_iter().map(|&v| v as f64).sum()
38}
39
40pub 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
59pub 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
86pub 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 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 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
144pub 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
154pub 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 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 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 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 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}