Skip to main content

qualia_core_db/platform/compute_bridge/
matrix.rs

1//! The per-class capability matrix and the built-in backends (plan §3).
2//!
3//! `benchmark_devices` (the existing AH-track probe) measures ONE GEMV and ranks
4//! circuits globally. This extends that to a **per-kernel-class** matrix: each class
5//! is measured on every available backend, and a winner is recorded *per class*
6//! (the fastest backend for GEMV is not the fastest for an FFT). The result is what
7//! `ComputePolicy::select` consults.
8//!
9//! Built-in backends registered by default:
10//! - [`CpuBackend`] — the native `rayon` reference kernels (`super::reference`),
11//!   always available, a row for every class.
12//! - [`WgpuBackend`] — the portable wgpu path. Today it measures the `DenseLinear`
13//!   class via the existing `device_benchmark` GEMV (a real GPU number); the other
14//!   classes have no portable GPU microkernel *yet*, so it returns no rows for them
15//!   (recorded honestly as "not probed on GPU", never faked). Those per-class GPU
16//!   kernels land per module (plan P2/P5) and appear here with zero further wiring.
17
18use std::time::Instant;
19
20use super::backend::{BackendId, KernelPanel, ProbeableBackend};
21use super::kernel_class::KernelClass;
22use super::reference;
23use crate::device_benchmark::{benchmark_devices, CircuitBench, CircuitKind};
24
25/// Measured per-class capability: for each class, the backend rows ranked
26/// fastest-first. `best_for` is the O(1)-ish lookup the STEM call sites use.
27#[derive(Debug, Clone)]
28pub struct ClassMatrix {
29    per_class: Vec<(KernelClass, Vec<CircuitBench>)>,
30}
31
32impl ClassMatrix {
33    /// Construct directly from per-class ranked rows — used by the passport loader
34    /// (a cached matrix) and by tests that synthesise a matrix without a GPU.
35    pub fn from_per_class(per_class: Vec<(KernelClass, Vec<CircuitBench>)>) -> Self {
36        Self { per_class }
37    }
38
39    /// Ranked rows for a class (fastest first); empty if the class was not probed.
40    pub fn rows(&self, class: KernelClass) -> &[CircuitBench] {
41        self.per_class
42            .iter()
43            .find(|(c, _)| *c == class)
44            .map(|(_, r)| r.as_slice())
45            .unwrap_or(&[])
46    }
47
48    /// The measured fastest (circuit, backend) for a class, if any was probed.
49    pub fn best_for(&self, class: KernelClass) -> Option<&CircuitBench> {
50        self.rows(class).first()
51    }
52
53    pub fn summary(&self) -> String {
54        let mut s = String::from("ClassMatrix (per kernel-class, ranked fastest-first):\n");
55        for (class, rows) in &self.per_class {
56            s.push_str(&format!("  {}:\n", class.label()));
57            for (i, r) in rows.iter().enumerate() {
58                s.push_str(&format!(
59                    "    {}. {:<26} [{:?}/{}] {:>9.4} ms  score {:.3}\n",
60                    i + 1,
61                    r.label,
62                    r.kind,
63                    r.backend,
64                    r.ms_per_gemv,
65                    r.rel_score
66                ));
67            }
68        }
69        s
70    }
71}
72
73/// Time `f` over `iters` runs (with one warmup) and return ms/run.
74fn time_ms(iters: u32, mut f: impl FnMut()) -> f64 {
75    f(); // warmup
76    let t0 = Instant::now();
77    for _ in 0..iters {
78        f();
79    }
80    t0.elapsed().as_secs_f64() * 1e3 / iters as f64
81}
82
83/// Native-CPU backend: the `rayon` reference kernels. Always available.
84pub struct CpuBackend;
85
86impl ProbeableBackend for CpuBackend {
87    fn id(&self) -> BackendId {
88        BackendId::CPU
89    }
90    fn available(&self) -> bool {
91        true // CPU is always present (plan §7)
92    }
93    fn probe_class(&self, class: KernelClass, panel: &KernelPanel) -> Vec<CircuitBench> {
94        let label = format!("CPU native (rayon, {} cores)", num_cpus::get());
95        let ms = match class {
96            KernelClass::DenseLinear => {
97                let n = panel.dense_n;
98                let w = vec![0.05f32; n * n];
99                let x = vec![0.1f32; n];
100                let mut y = vec![0.0f32; n];
101                time_ms(5, || reference::gemv(&w, &x, &mut y))
102            }
103            KernelClass::ElementwiseMap => {
104                let x = vec![0.1f32; panel.vector_len];
105                let mut y = vec![0.0f32; panel.vector_len];
106                time_ms(10, || reference::axpb(2.0, &x, 1.0, &mut y))
107            }
108            KernelClass::Reduction => {
109                let x = vec![0.1f32; panel.vector_len];
110                time_ms(10, || {
111                    let _ = reference::reduce_sum(&x);
112                })
113            }
114            KernelClass::Stencil => {
115                let x = vec![0.1f32; panel.grid_n];
116                let mut y = vec![0.0f32; panel.grid_n];
117                time_ms(10, || reference::stencil3(&x, &mut y))
118            }
119            KernelClass::AllPairs => {
120                let pts: Vec<f32> = (0..panel.nbody_n * 3)
121                    .map(|i| (i % 97) as f32 * 0.1)
122                    .collect();
123                time_ms(3, || {
124                    let _ = reference::allpairs_potential(&pts);
125                })
126            }
127            KernelClass::Fft => {
128                let n = panel.fft_n.next_power_of_two();
129                let base: Vec<f32> = (0..n).map(|i| (i as f32 * 0.01).sin()).collect();
130                time_ms(5, || {
131                    let mut re = base.clone();
132                    let mut im = vec![0.0f32; n];
133                    reference::fft_radix2(&mut re, &mut im, false);
134                })
135            }
136            KernelClass::Scan => {
137                let x = vec![0.1f32; panel.vector_len];
138                let mut y = vec![0.0f32; panel.vector_len];
139                time_ms(10, || reference::prefix_sum(&x, &mut y))
140            }
141            KernelClass::Divergent => time_ms(3, || {
142                let _ = reference::monte_carlo_pi(panel.mc_steps);
143            }),
144        };
145        vec![CircuitBench {
146            label,
147            kind: CircuitKind::Cpu,
148            backend: "native".to_string(),
149            ms_per_gemv: ms,
150            gflops: 0.0, // throughput proxy not derived here; ranking uses ms
151            upload_gbps: f64::INFINITY, // data already in the CPU pool — no transfer
152            rel_score: 1.0,
153            decode_proxy_tok_s: None,
154        }]
155    }
156}
157
158/// Portable wgpu backend. Measures `DenseLinear` via the existing GEMV probe; other
159/// classes have no portable GPU microkernel yet (returns no rows — honest).
160pub struct WgpuBackend;
161
162impl ProbeableBackend for WgpuBackend {
163    fn id(&self) -> BackendId {
164        BackendId::WGPU
165    }
166    fn available(&self) -> bool {
167        // Available if any non-CPU wgpu adapter enumerates. Cheap-ish; only at boot.
168        let instance = wgpu::Instance::default();
169        pollster::block_on(instance.enumerate_adapters(wgpu::Backends::all()))
170            .iter()
171            .any(|a| {
172                let info = a.get_info();
173                info.device_type != wgpu::DeviceType::Cpu && info.device != 0
174            })
175    }
176    fn probe_class(&self, class: KernelClass, panel: &KernelPanel) -> Vec<CircuitBench> {
177        match class {
178            KernelClass::DenseLinear => {
179                // Reuse the real GEMV probe; keep only the GPU rows (the CPU row is
180                // contributed by CpuBackend).
181                benchmark_devices(panel.dense_n)
182                    .circuits
183                    .into_iter()
184                    .filter(|c| c.kind != CircuitKind::Cpu)
185                    .collect()
186            }
187            // No portable GPU microkernel for these classes yet → no rows (plan P2/P5).
188            _ => Vec::new(),
189        }
190    }
191}
192
193/// Probe every available backend across every kernel class and assemble the ranked
194/// per-class matrix. Boot-time only (cache it in the passport).
195pub fn probe_class_matrix(
196    registry: &super::backend::BackendRegistry,
197    panel: &KernelPanel,
198) -> ClassMatrix {
199    let mut per_class = Vec::with_capacity(KernelClass::ALL.len());
200    for class in KernelClass::ALL {
201        let mut rows: Vec<CircuitBench> = Vec::new();
202        for backend in registry.available() {
203            rows.extend(backend.probe_class(class, panel));
204        }
205        // Rank fastest-first by measured ms, fill relative scores.
206        rows.sort_by(|a, b| {
207            a.ms_per_gemv
208                .partial_cmp(&b.ms_per_gemv)
209                .unwrap_or(std::cmp::Ordering::Equal)
210        });
211        if let Some(best) = rows.first().map(|c| c.ms_per_gemv) {
212            for r in &mut rows {
213                r.rel_score = if r.ms_per_gemv > 0.0 {
214                    best / r.ms_per_gemv
215                } else {
216                    0.0
217                };
218            }
219        }
220        per_class.push((class, rows));
221    }
222    ClassMatrix { per_class }
223}
224
225#[cfg(test)]
226mod tests {
227    use super::super::backend::BackendRegistry;
228    use super::*;
229
230    #[test]
231    fn cpu_backend_probes_every_class_with_real_rows() {
232        let cpu = CpuBackend;
233        let panel = KernelPanel::quick();
234        for class in KernelClass::ALL {
235            let rows = cpu.probe_class(class, &panel);
236            assert_eq!(rows.len(), 1, "CPU must yield a row for {}", class.label());
237            assert!(
238                rows[0].ms_per_gemv > 0.0,
239                "{} CPU time must be positive",
240                class.label()
241            );
242            assert!(rows[0].upload_gbps.is_infinite(), "CPU is in-pool");
243        }
244    }
245
246    #[test]
247    fn class_matrix_has_cpu_winner_for_every_class_headless() {
248        // CPU-only registry (deterministic, no GPU): every class gets a ranked row,
249        // and the best is score 1.0.
250        let mut reg = BackendRegistry::new();
251        reg.register(Box::new(CpuBackend));
252        let m = probe_class_matrix(&reg, &KernelPanel::quick());
253        for class in KernelClass::ALL {
254            let best = m.best_for(class).expect("a class winner must exist");
255            assert_eq!(best.kind, CircuitKind::Cpu);
256            assert!((best.rel_score - 1.0).abs() < 1e-9);
257        }
258    }
259}