Skip to main content

cortiq_engine/
f32_backend.rs

1//! Explicit, thread-local F32 execution policy. The historical scalar path is
2//! the default. Optimized reductions are NOT bit-exact: callers must audit the
3//! complete model before opting in. This changes runtime arithmetic, not CMF.
4use std::cell::Cell;
5
6#[derive(Clone, Copy, Debug, Default, PartialEq, Eq)]
7pub enum Backend {
8    #[default]
9    Reference,
10    Accelerate,
11}
12
13thread_local! {
14    static BACKEND: Cell<Backend> = const { Cell::new(Backend::Reference) };
15    static CALLS: Cell<u64> = const { Cell::new(0) };
16}
17
18pub fn available(backend: Backend) -> bool {
19    backend == Backend::Reference || cfg!(target_os = "macos")
20}
21
22pub fn active() -> bool {
23    BACKEND.with(|b| b.get() != Backend::Reference)
24}
25
26pub fn calls() -> u64 {
27    CALLS.with(Cell::get)
28}
29
30/// Nested scopes and unwinding restore the previous policy; other threads are
31/// unaffected. An unavailable explicit backend is an error, never a false label.
32pub fn scope<T>(backend: Backend, f: impl FnOnce() -> T) -> T {
33    assert!(available(backend), "requested F32 backend is unavailable");
34    struct Restore(Backend);
35    impl Drop for Restore {
36        fn drop(&mut self) {
37            BACKEND.with(|b| b.set(self.0));
38        }
39    }
40    let _restore = Restore(BACKEND.with(|b| b.replace(backend)));
41    f()
42}
43
44#[cfg(target_os = "macos")]
45#[link(name = "Accelerate", kind = "framework")]
46unsafe extern "C" {
47    fn cblas_sgemv(
48        order: i32,
49        trans: i32,
50        m: i32,
51        n: i32,
52        alpha: f32,
53        a: *const f32,
54        lda: i32,
55        x: *const f32,
56        incx: i32,
57        beta: f32,
58        y: *mut f32,
59        incy: i32,
60    );
61    fn cblas_sgemm(
62        order: i32,
63        ta: i32,
64        tb: i32,
65        m: i32,
66        n: i32,
67        k: i32,
68        alpha: f32,
69        a: *const f32,
70        lda: i32,
71        b: *const f32,
72        ldb: i32,
73        beta: f32,
74        c: *mut f32,
75        ldc: i32,
76    );
77}
78
79/// Historical F32 matvec permits a short output and uses x.len() as stride.
80pub(crate) fn matvec(w: &[f32], x: &[f32], y: &mut [f32]) -> bool {
81    if !active() {
82        return false;
83    }
84    #[cfg(target_os = "macos")]
85    {
86        let (m, n) = (y.len(), x.len());
87        assert!(
88            m.checked_mul(n).is_some_and(|len| len <= w.len()),
89            "short F32 weights"
90        );
91        if m == 0 {
92            return true;
93        }
94        if n == 0 {
95            y.fill(0.);
96            return true;
97        }
98        let (Ok(m), Ok(n)) = (i32::try_from(m), i32::try_from(n)) else {
99            return false;
100        };
101        // RowMajor=101, NoTrans=111, beta=0: all output cells are overwritten.
102        unsafe {
103            cblas_sgemv(
104                101,
105                111,
106                m,
107                n,
108                1.,
109                w.as_ptr(),
110                n,
111                x.as_ptr(),
112                1,
113                0.,
114                y.as_mut_ptr(),
115                1,
116            );
117        }
118        CALLS.with(|c| c.set(c.get() + 1));
119        return true;
120    }
121    #[allow(unreachable_code)]
122    false
123}
124
125/// Input [batch, cols], weights [rows, cols], output [batch, rows].
126pub(crate) fn matmat(
127    w: &[f32],
128    x: &[f32],
129    batch: usize,
130    rows: usize,
131    cols: usize,
132    y: &mut [f32],
133) -> bool {
134    if !active() {
135        return false;
136    }
137    #[cfg(target_os = "macos")]
138    {
139        assert_eq!(batch.checked_mul(cols), Some(x.len()), "F32 input shape");
140        assert_eq!(batch.checked_mul(rows), Some(y.len()), "F32 output shape");
141        assert_eq!(rows.checked_mul(cols), Some(w.len()), "F32 weight shape");
142        if batch == 0 || rows == 0 {
143            return true;
144        }
145        if cols == 0 {
146            y.fill(0.);
147            return true;
148        }
149        let (Ok(m), Ok(n), Ok(k)) = (
150            i32::try_from(batch),
151            i32::try_from(rows),
152            i32::try_from(cols),
153        ) else {
154            return false;
155        };
156        unsafe {
157            cblas_sgemm(
158                101,
159                111,
160                112,
161                m,
162                n,
163                k,
164                1.,
165                x.as_ptr(),
166                k,
167                w.as_ptr(),
168                k,
169                0.,
170                y.as_mut_ptr(),
171                n,
172            );
173        }
174        CALLS.with(|c| c.set(c.get() + 1));
175        return true;
176    }
177    #[allow(unreachable_code)]
178    false
179}
180
181#[cfg(test)]
182mod tests {
183    use super::*;
184    #[test]
185    fn reference_is_default_and_does_not_write() {
186        let mut y = [7.];
187        assert!(!matvec(&[2.], &[3.], &mut y));
188        assert_eq!(y, [7.]);
189    }
190    #[test]
191    #[cfg(target_os = "macos")]
192    fn scoped_blas_shapes_overwrite_and_restore() {
193        scope(Backend::Accelerate, || {
194            let w = [1., 2., 3., -1., 4., 2.];
195            let x = [2., 3., 4.];
196            let mut y = [f32::NAN; 2];
197            assert!(matvec(&w, &x, &mut y));
198            assert_eq!(y, [20., 18.]);
199            assert!(matvec(&w, &x, &mut y[..1]));
200            let mut ys = [f32::NAN; 4];
201            assert!(matmat(&w, &[2., 3., 4., 1., 0., -1.], 2, 2, 3, &mut ys));
202            assert_eq!(ys, [20., 18., -2., -3.]);
203            scope(Backend::Reference, || assert!(!active()));
204            assert!(active());
205            std::thread::spawn(|| assert!(!active())).join().unwrap();
206        });
207        assert!(!active());
208        let _ = std::panic::catch_unwind(|| scope(Backend::Accelerate, || panic!("restore")));
209        assert!(!active());
210    }
211    #[test]
212    #[cfg(target_os = "macos")]
213    fn qtensor_single_many_and_batch_match_reference() {
214        use crate::{pool::Pool, qtensor::QTensor};
215        let pool = Pool::new(2);
216        for (rows, cols, batch) in [(1, 1, 1), (5, 7, 3), (300, 64, 7)] {
217            let w: Vec<_> = (0..rows * cols)
218                .map(|i| ((i * 7 % 97) as f32 - 48.) / 49.)
219                .collect();
220            let t = QTensor::from_f32(w, rows, cols);
221            let xs: Vec<_> = (0..batch * cols).map(|i| (i as f32 * 0.17).sin()).collect();
222            let mut reference = vec![0.; batch * rows];
223            t.matmat(&xs, batch, &mut reference, Some(&pool));
224            scope(Backend::Accelerate, || {
225                let mut got = vec![f32::NAN; batch * rows];
226                t.matmat(&xs, batch, &mut got, Some(&pool));
227                for (a, b) in got.iter().zip(&reference) {
228                    assert!((a - b).abs() < 2e-5);
229                }
230                let (mut a, mut b) = (vec![f32::NAN; rows], vec![f32::NAN; rows]);
231                QTensor::matvec_many([&t, &t], &xs[..cols], [&mut a, &mut b], Some(&pool));
232                assert_eq!(a, b);
233                for (a, b) in a.iter().zip(&reference) {
234                    assert!((a - b).abs() < 2e-5);
235                }
236            });
237        }
238    }
239    #[test]
240    #[cfg(target_os = "macos")]
241    fn empty_and_malformed_shapes_never_reach_ffi() {
242        scope(Backend::Accelerate, || {
243            let mut y = [f32::NAN; 3];
244            assert!(matvec(&[], &[], &mut y));
245            assert_eq!(y, [0.; 3]);
246            assert!(matmat(&[], &[], 1, 3, 0, &mut y));
247            assert_eq!(y, [0.; 3]);
248            assert!(std::panic::catch_unwind(|| matvec(&[1.], &[1., 2.], &mut [0.])).is_err());
249            assert!(
250                std::panic::catch_unwind(|| matmat(&[1.], &[1.], 2, 1, 1, &mut [0.; 2])).is_err()
251            );
252        });
253    }
254}