1use 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
30pub 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
79pub(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 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
125pub(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}