sparse_ir_core/gemm.rs
1//! Column-major GEMM with a pluggable BLAS backend.
2//!
3//! All matrices are column-major (Fortran/BLAS convention) and every call is
4//! expressed in BLAS terms: `C <- alpha * op(A) * op(B) + beta * C` with
5//! explicit leading dimensions, so strided sub-blocks are addressed without
6//! copies.
7//!
8//! # Design
9//! - **Default**: pure Rust faer backend (sequential), or system BLAS when
10//! the `system-blas` feature is enabled.
11//! - **Optional**: external LP64/ILP64 BLAS injected through function
12//! pointers (used by the C API).
13//! - **Thread-safe**: the process-wide default is protected by an `RwLock`;
14//! per-call backends are passed as [`GemmBackendHandle`].
15//!
16//! The safe entry point is [`gemm`], which validates buffer lengths and
17//! leading dimensions before dispatching; [`matmul`] multiplies two
18//! [`Matrix`](crate::Matrix) values.
19//!
20//! # Example
21//! ```
22//! use sparse_ir::Matrix;
23//! use sparse_ir::gemm::{GemmBackendHandle, matmul};
24//!
25//! // Column-major data: [[1, 2], [3, 4]] and [[5, 6], [7, 8]]
26//! let a = Matrix::<f64>::from_vec_col_major([2, 2], vec![1.0, 3.0, 2.0, 4.0]).unwrap();
27//! let b = Matrix::<f64>::from_vec_col_major([2, 2], vec![5.0, 7.0, 6.0, 8.0]).unwrap();
28//!
29//! // `None` uses the global dispatcher: the Faer backend by default (system
30//! // BLAS with the `system-blas` feature, or a BLAS injected at runtime)
31//! let c = matmul(None, &a, &b).unwrap();
32//! assert_eq!(c.host_data().unwrap(), &[19.0, 43.0, 22.0, 50.0]); // [[19, 22], [43, 50]]
33//!
34//! // Or pass an explicit backend handle instead of relying on global state
35//! let faer = GemmBackendHandle::default();
36//! assert_eq!(
37//! matmul(Some(&faer), &a, &b).unwrap().host_data().unwrap(),
38//! c.host_data().unwrap()
39//! );
40//! ```
41//!
42//! A custom BLAS (e.g. from the C API) is injected at runtime with
43//! [`set_blas_backend`] (LP64) or [`set_ilp64_backend`] (ILP64); see their
44//! examples. [`clear_blas_backend`] restores the Faer backend.
45
46use num_complex::Complex;
47use num_traits::{One, Zero};
48use once_cell::sync::Lazy;
49use std::sync::{Arc, RwLock};
50
51//==============================================================================
52// BLAS Function Pointer Types
53//==============================================================================
54
55/// BLAS dgemm function pointer type (LP64: 32-bit integers)
56///
57/// Signature matches Fortran BLAS dgemm:
58/// ```c
59/// void dgemm_(char *transa, char *transb, int *m, int *n, int *k,
60/// double *alpha, double *a, int *lda, double *b, int *ldb,
61/// double *beta, double *c, int *ldc);
62/// ```
63/// Note: All parameters are passed by reference (pointers).
64/// Transpose options: 'N' (no transpose), 'T' (transpose), 'C' (conjugate transpose).
65pub type DgemmFnPtr = unsafe extern "C" fn(
66 transa: *const libc::c_char,
67 transb: *const libc::c_char,
68 m: *const libc::c_int,
69 n: *const libc::c_int,
70 k: *const libc::c_int,
71 alpha: *const libc::c_double,
72 a: *const libc::c_double,
73 lda: *const libc::c_int,
74 b: *const libc::c_double,
75 ldb: *const libc::c_int,
76 beta: *const libc::c_double,
77 c: *mut libc::c_double,
78 ldc: *const libc::c_int,
79);
80
81/// BLAS zgemm function pointer type (LP64: 32-bit integers)
82///
83/// Signature matches Fortran BLAS zgemm:
84/// ```c
85/// void zgemm_(char *transa, char *transb, int *m, int *n, int *k,
86/// void *alpha, void *a, int *lda, void *b, int *ldb,
87/// void *beta, void *c, int *ldc);
88/// ```
89/// Note: All parameters are passed by reference (pointers).
90/// Complex numbers are passed as void* (typically complex<double>*).
91/// Transpose options: 'N' (no transpose), 'T' (transpose), 'C' (conjugate transpose).
92pub type ZgemmFnPtr = unsafe extern "C" fn(
93 transa: *const libc::c_char,
94 transb: *const libc::c_char,
95 m: *const libc::c_int,
96 n: *const libc::c_int,
97 k: *const libc::c_int,
98 alpha: *const num_complex::Complex<f64>,
99 a: *const num_complex::Complex<f64>,
100 lda: *const libc::c_int,
101 b: *const num_complex::Complex<f64>,
102 ldb: *const libc::c_int,
103 beta: *const num_complex::Complex<f64>,
104 c: *mut num_complex::Complex<f64>,
105 ldc: *const libc::c_int,
106);
107
108// When using system BLAS via `blas-sys`, we need a small wrapper to adapt
109// `blas_sys::zgemm_` (which uses `c_double_complex = [f64; 2]`) to the
110// `ZgemmFnPtr` signature that takes `num_complex::Complex<f64>`.
111#[cfg(feature = "system-blas")]
112unsafe extern "C" fn zgemm_wrapper(
113 transa: *const libc::c_char,
114 transb: *const libc::c_char,
115 m: *const libc::c_int,
116 n: *const libc::c_int,
117 k: *const libc::c_int,
118 alpha: *const num_complex::Complex<f64>,
119 a: *const num_complex::Complex<f64>,
120 lda: *const libc::c_int,
121 b: *const num_complex::Complex<f64>,
122 ldb: *const libc::c_int,
123 beta: *const num_complex::Complex<f64>,
124 c: *mut num_complex::Complex<f64>,
125 ldc: *const libc::c_int,
126) {
127 // Safety: `blas_sys::c_double_complex` is defined as `[f64; 2]` and is
128 // layout-compatible with `num_complex::Complex<f64>` in memory, so we can
129 // cast between the two pointer types here.
130 unsafe {
131 blas_sys::zgemm_(
132 transa,
133 transb,
134 m,
135 n,
136 k,
137 alpha as *const _ as *const blas_sys::c_double_complex,
138 a as *const _ as *const blas_sys::c_double_complex,
139 lda,
140 b as *const _ as *const blas_sys::c_double_complex,
141 ldb,
142 beta as *const _ as *const blas_sys::c_double_complex,
143 c as *mut _ as *mut blas_sys::c_double_complex,
144 ldc,
145 );
146 }
147}
148
149/// BLAS dgemm function pointer type (ILP64: 64-bit integers)
150///
151/// Signature matches Fortran BLAS dgemm (ILP64):
152/// ```c
153/// void dgemm_(char *transa, char *transb, long long *m, long long *n, long long *k,
154/// double *alpha, double *a, long long *lda, double *b, long long *ldb,
155/// double *beta, double *c, long long *ldc);
156/// ```
157pub type Dgemm64FnPtr = unsafe extern "C" fn(
158 transa: *const libc::c_char,
159 transb: *const libc::c_char,
160 m: *const i64,
161 n: *const i64,
162 k: *const i64,
163 alpha: *const libc::c_double,
164 a: *const libc::c_double,
165 lda: *const i64,
166 b: *const libc::c_double,
167 ldb: *const i64,
168 beta: *const libc::c_double,
169 c: *mut libc::c_double,
170 ldc: *const i64,
171);
172
173/// BLAS zgemm function pointer type (ILP64: 64-bit integers)
174///
175/// Signature matches Fortran BLAS zgemm (ILP64):
176/// ```c
177/// void zgemm_(char *transa, char *transb, long long *m, long long *n, long long *k,
178/// void *alpha, void *a, long long *lda, void *b, long long *ldb,
179/// void *beta, void *c, long long *ldc);
180/// ```
181pub type Zgemm64FnPtr = unsafe extern "C" fn(
182 transa: *const libc::c_char,
183 transb: *const libc::c_char,
184 m: *const i64,
185 n: *const i64,
186 k: *const i64,
187 alpha: *const num_complex::Complex<f64>,
188 a: *const num_complex::Complex<f64>,
189 lda: *const i64,
190 b: *const num_complex::Complex<f64>,
191 ldb: *const i64,
192 beta: *const num_complex::Complex<f64>,
193 c: *mut num_complex::Complex<f64>,
194 ldc: *const i64,
195);
196
197//==============================================================================
198// Operation descriptors and errors
199//==============================================================================
200
201/// Operation applied to a GEMM operand.
202#[derive(Clone, Copy, Debug, PartialEq, Eq)]
203pub enum Transpose {
204 /// `op(X) = X`
205 N,
206 /// `op(X) = X^T`
207 T,
208 /// `op(X) = X^H`
209 C,
210}
211
212impl Transpose {
213 fn as_blas_char(self) -> libc::c_char {
214 (match self {
215 Transpose::N => b'N',
216 Transpose::T => b'T',
217 Transpose::C => b'C',
218 }) as libc::c_char
219 }
220}
221
222/// Errors reported by GEMM dispatch.
223#[derive(Clone, Debug, PartialEq, Eq, thiserror::Error)]
224pub enum GemmError {
225 /// A dimension or leading dimension does not fit the backend integer type.
226 #[error("GEMM argument {name}={value} exceeds the {abi} integer range")]
227 DimensionOverflow {
228 name: &'static str,
229 value: usize,
230 abi: &'static str,
231 },
232 /// A buffer is too short or a leading dimension is too small.
233 #[error("invalid GEMM argument: {0}")]
234 InvalidArgument(String),
235}
236
237//==============================================================================
238// GemmBackend Trait
239//==============================================================================
240
241/// GEMM backend trait for runtime dispatch.
242///
243/// Implementations compute `C <- alpha * op(A) * op(B) + beta * C` on
244/// column-major operands, with BLAS semantics: `op(A)` is `m x k`, `op(B)`
245/// is `k x n`, `C` is `m x n`, and `C` is not read when `beta == 0`.
246pub trait GemmBackend: Send + Sync {
247 /// Real double-precision GEMM.
248 ///
249 /// # Safety
250 /// The pointers must address column-major matrices of the stated shapes
251 /// and leading dimensions (`lda >= max(1, rows(A))`, and likewise for `B`
252 /// and `C`); `c` must not alias `a` or `b`.
253 #[allow(clippy::too_many_arguments)]
254 unsafe fn dgemm(
255 &self,
256 transa: Transpose,
257 transb: Transpose,
258 m: usize,
259 n: usize,
260 k: usize,
261 alpha: f64,
262 a: *const f64,
263 lda: usize,
264 b: *const f64,
265 ldb: usize,
266 beta: f64,
267 c: *mut f64,
268 ldc: usize,
269 ) -> Result<(), GemmError>;
270
271 /// Complex double-precision GEMM.
272 ///
273 /// # Safety
274 /// Same contract as [`GemmBackend::dgemm`].
275 #[allow(clippy::too_many_arguments)]
276 unsafe fn zgemm(
277 &self,
278 transa: Transpose,
279 transb: Transpose,
280 m: usize,
281 n: usize,
282 k: usize,
283 alpha: Complex<f64>,
284 a: *const Complex<f64>,
285 lda: usize,
286 b: *const Complex<f64>,
287 ldb: usize,
288 beta: Complex<f64>,
289 c: *mut Complex<f64>,
290 ldc: usize,
291 ) -> Result<(), GemmError>;
292
293 /// Returns true if this backend uses 64-bit integers (ILP64).
294 fn is_ilp64(&self) -> bool {
295 false
296 }
297
298 /// Backend name for diagnostics.
299 fn name(&self) -> &'static str;
300}
301
302//==============================================================================
303// Faer Backend (Default, Pure Rust, Zero-Copy)
304//==============================================================================
305
306/// Pure Rust faer backend (sequential).
307struct FaerBackend;
308
309/// Shared faer implementation for `f64` and `Complex<f64>`.
310///
311/// # Safety
312/// Same contract as [`GemmBackend::dgemm`].
313#[allow(clippy::too_many_arguments)]
314unsafe fn faer_gemm<T>(
315 transa: Transpose,
316 transb: Transpose,
317 m: usize,
318 n: usize,
319 k: usize,
320 alpha: T,
321 a: *const T,
322 lda: usize,
323 b: *const T,
324 ldb: usize,
325 beta: T,
326 c: *mut T,
327 ldc: usize,
328) where
329 T: faer_traits::ComplexField + Copy + PartialEq + Zero + One + std::ops::MulAssign,
330{
331 use faer::linalg::matmul::matmul_with_conj;
332 use faer::mat::{MatMut, MatRef};
333 use faer::{Accum, Conj, Par};
334
335 if m == 0 || n == 0 {
336 return;
337 }
338 // SAFETY: the caller guarantees `c` addresses an m x n column-major
339 // matrix with leading dimension ldc >= m; `c` is non-null because m, n > 0.
340 let mut dst = unsafe { MatMut::from_raw_parts_mut(c, m, n, 1isize, ldc as isize) };
341
342 // faer accumulates with beta in {0, 1}; other values scale C first.
343 let accum = if beta == T::zero() {
344 Accum::Replace
345 } else {
346 if beta != T::one() {
347 for j in 0..n {
348 for i in 0..m {
349 // SAFETY: (i, j) is within the m x n destination.
350 unsafe { *dst.as_mut().get_mut_unchecked(i, j) *= beta };
351 }
352 }
353 }
354 Accum::Add
355 };
356 if k == 0 {
357 if accum == Accum::Replace {
358 dst.fill(T::zero());
359 }
360 return;
361 }
362
363 let operand = |ptr: *const T, op: Transpose, rows: usize, cols: usize, ld: usize| {
364 // SAFETY: the caller guarantees `ptr` addresses the stored (untransposed)
365 // `rows x cols` operand with leading dimension ld; rows, cols > 0.
366 let stored = match op {
367 Transpose::N => unsafe { MatRef::from_raw_parts(ptr, rows, cols, 1isize, ld as isize) },
368 Transpose::T | Transpose::C => {
369 unsafe { MatRef::from_raw_parts(ptr, cols, rows, 1isize, ld as isize) }.transpose()
370 }
371 };
372 let conj = if op == Transpose::C {
373 Conj::Yes
374 } else {
375 Conj::No
376 };
377 (stored, conj)
378 };
379 let (lhs, conj_lhs) = operand(a, transa, m, k, lda);
380 let (rhs, conj_rhs) = operand(b, transb, k, n, ldb);
381 matmul_with_conj(dst, accum, lhs, conj_lhs, rhs, conj_rhs, alpha, Par::Seq);
382}
383
384impl GemmBackend for FaerBackend {
385 unsafe fn dgemm(
386 &self,
387 transa: Transpose,
388 transb: Transpose,
389 m: usize,
390 n: usize,
391 k: usize,
392 alpha: f64,
393 a: *const f64,
394 lda: usize,
395 b: *const f64,
396 ldb: usize,
397 beta: f64,
398 c: *mut f64,
399 ldc: usize,
400 ) -> Result<(), GemmError> {
401 // SAFETY: forwarded caller contract.
402 unsafe { faer_gemm(transa, transb, m, n, k, alpha, a, lda, b, ldb, beta, c, ldc) };
403 Ok(())
404 }
405
406 unsafe fn zgemm(
407 &self,
408 transa: Transpose,
409 transb: Transpose,
410 m: usize,
411 n: usize,
412 k: usize,
413 alpha: Complex<f64>,
414 a: *const Complex<f64>,
415 lda: usize,
416 b: *const Complex<f64>,
417 ldb: usize,
418 beta: Complex<f64>,
419 c: *mut Complex<f64>,
420 ldc: usize,
421 ) -> Result<(), GemmError> {
422 // SAFETY: forwarded caller contract.
423 unsafe { faer_gemm(transa, transb, m, n, k, alpha, a, lda, b, ldb, beta, c, ldc) };
424 Ok(())
425 }
426
427 fn name(&self) -> &'static str {
428 "Faer (Pure Rust)"
429 }
430}
431
432//==============================================================================
433// External BLAS Backends (LP64 and ILP64)
434//==============================================================================
435
436/// Checked conversion of the six BLAS integer arguments.
437fn blas_ints<I: TryFrom<usize>>(
438 abi: &'static str,
439 m: usize,
440 n: usize,
441 k: usize,
442 lda: usize,
443 ldb: usize,
444 ldc: usize,
445) -> Result<[I; 6], GemmError> {
446 let conv = |name: &'static str, value: usize| {
447 I::try_from(value).map_err(|_| GemmError::DimensionOverflow { name, value, abi })
448 };
449 Ok([
450 conv("m", m)?,
451 conv("n", n)?,
452 conv("k", k)?,
453 conv("lda", lda)?,
454 conv("ldb", ldb)?,
455 conv("ldc", ldc)?,
456 ])
457}
458
459/// External BLAS backend (LP64: 32-bit integers).
460pub struct ExternalBlasBackend {
461 dgemm: DgemmFnPtr,
462 zgemm: ZgemmFnPtr,
463}
464
465impl ExternalBlasBackend {
466 /// Wrap LP64 Fortran `dgemm_`/`zgemm_` function pointers.
467 pub fn new(dgemm: DgemmFnPtr, zgemm: ZgemmFnPtr) -> Self {
468 Self { dgemm, zgemm }
469 }
470}
471
472/// External BLAS backend (ILP64: 64-bit integers).
473pub struct ExternalBlas64Backend {
474 dgemm64: Dgemm64FnPtr,
475 zgemm64: Zgemm64FnPtr,
476}
477
478impl ExternalBlas64Backend {
479 /// Wrap ILP64 Fortran `dgemm_`/`zgemm_` function pointers.
480 pub fn new(dgemm64: Dgemm64FnPtr, zgemm64: Zgemm64FnPtr) -> Self {
481 Self { dgemm64, zgemm64 }
482 }
483}
484
485macro_rules! impl_external_backend {
486 ($ty:ty, $int:ty, $abi:literal, $dfield:ident, $zfield:ident, $ilp64:expr, $name:literal) => {
487 impl GemmBackend for $ty {
488 unsafe fn dgemm(
489 &self,
490 transa: Transpose,
491 transb: Transpose,
492 m: usize,
493 n: usize,
494 k: usize,
495 alpha: f64,
496 a: *const f64,
497 lda: usize,
498 b: *const f64,
499 ldb: usize,
500 beta: f64,
501 c: *mut f64,
502 ldc: usize,
503 ) -> Result<(), GemmError> {
504 if m == 0 || n == 0 {
505 return Ok(());
506 }
507 let [m, n, k, lda, ldb, ldc] = blas_ints::<$int>($abi, m, n, k, lda, ldb, ldc)?;
508 let (ta, tb) = (transa.as_blas_char(), transb.as_blas_char());
509 // SAFETY: arguments follow the Fortran BLAS reference-passing
510 // convention; the caller guarantees the operand contract.
511 unsafe {
512 (self.$dfield)(
513 &ta, &tb, &m, &n, &k, &alpha, a, &lda, b, &ldb, &beta, c, &ldc,
514 )
515 };
516 Ok(())
517 }
518
519 unsafe fn zgemm(
520 &self,
521 transa: Transpose,
522 transb: Transpose,
523 m: usize,
524 n: usize,
525 k: usize,
526 alpha: Complex<f64>,
527 a: *const Complex<f64>,
528 lda: usize,
529 b: *const Complex<f64>,
530 ldb: usize,
531 beta: Complex<f64>,
532 c: *mut Complex<f64>,
533 ldc: usize,
534 ) -> Result<(), GemmError> {
535 if m == 0 || n == 0 {
536 return Ok(());
537 }
538 let [m, n, k, lda, ldb, ldc] = blas_ints::<$int>($abi, m, n, k, lda, ldb, ldc)?;
539 let (ta, tb) = (transa.as_blas_char(), transb.as_blas_char());
540 // SAFETY: as for dgemm; Complex<f64> is layout-compatible with
541 // Fortran COMPLEX*16.
542 unsafe {
543 (self.$zfield)(
544 &ta, &tb, &m, &n, &k, &alpha, a, &lda, b, &ldb, &beta, c, &ldc,
545 )
546 };
547 Ok(())
548 }
549
550 fn is_ilp64(&self) -> bool {
551 $ilp64
552 }
553
554 fn name(&self) -> &'static str {
555 $name
556 }
557 }
558 };
559}
560
561impl_external_backend!(
562 ExternalBlasBackend,
563 libc::c_int,
564 "LP64",
565 dgemm,
566 zgemm,
567 false,
568 "External BLAS (LP64)"
569);
570impl_external_backend!(
571 ExternalBlas64Backend,
572 i64,
573 "ILP64",
574 dgemm64,
575 zgemm64,
576 true,
577 "External BLAS (ILP64)"
578);
579
580//==============================================================================
581// Backend Handle
582//==============================================================================
583
584/// Shared, cloneable GEMM backend selection.
585///
586/// # Example
587/// ```
588/// use sparse_ir::Matrix;
589/// use sparse_ir::gemm::{GemmBackendHandle, matmul};
590///
591/// // Column-major data: [[1, 2], [3, 4]] and [[5, 6], [7, 8]]
592/// let a = Matrix::<f64>::from_vec_col_major([2, 2], vec![1.0, 3.0, 2.0, 4.0]).unwrap();
593/// let b = Matrix::<f64>::from_vec_col_major([2, 2], vec![5.0, 7.0, 6.0, 8.0]).unwrap();
594/// let backend = GemmBackendHandle::default();
595/// let result = matmul(Some(&backend), &a, &b).unwrap();
596/// assert_eq!(result.host_data().unwrap(), &[19.0, 43.0, 22.0, 50.0]); // [[19, 22], [43, 50]]
597/// ```
598#[derive(Clone)]
599pub struct GemmBackendHandle {
600 inner: Arc<dyn GemmBackend>,
601}
602
603impl GemmBackendHandle {
604 /// Wrap a backend.
605 pub fn new(backend: Box<dyn GemmBackend>) -> Self {
606 Self {
607 inner: Arc::from(backend),
608 }
609 }
610
611 /// Pure Rust faer backend.
612 #[allow(clippy::should_implement_trait)]
613 pub fn default() -> Self {
614 Self {
615 inner: Arc::new(FaerBackend),
616 }
617 }
618
619 pub(crate) fn as_ref(&self) -> &dyn GemmBackend {
620 self.inner.as_ref()
621 }
622}
623
624//==============================================================================
625// Global Dispatcher
626//==============================================================================
627
628static BLAS_DISPATCHER: Lazy<RwLock<Box<dyn GemmBackend>>> = Lazy::new(|| {
629 #[cfg(feature = "system-blas")]
630 {
631 // Use system BLAS (LP64) by default via `blas-sys`.
632 let backend =
633 ExternalBlasBackend::new(blas_sys::dgemm_ as DgemmFnPtr, zgemm_wrapper as ZgemmFnPtr);
634 RwLock::new(Box::new(backend) as Box<dyn GemmBackend>)
635 }
636 #[cfg(not(feature = "system-blas"))]
637 {
638 RwLock::new(Box::new(FaerBackend) as Box<dyn GemmBackend>)
639 }
640});
641
642/// Set the process-wide default to an LP64 BLAS.
643///
644/// # Safety
645/// The function pointers must be valid Fortran-convention LP64 `dgemm_` and
646/// `zgemm_` implementations that stay valid for the rest of the process.
647///
648/// # Example
649/// ```
650/// use num_complex::Complex;
651/// use sparse_ir::Matrix;
652/// use sparse_ir::gemm::{clear_blas_backend, get_backend_info, matmul, set_blas_backend};
653/// # use std::ffi::{c_char, c_int};
654/// # use std::ops::{Add, Mul};
655/// #
656/// # // Naive column-major C = alpha * A * B + beta * C for transa = transb = 'N'
657/// # // (what sparse-ir passes), standing in for a BLAS library.
658/// # unsafe fn gemm_nn<T>(mnk: [usize; 3], alpha: T, a: *const T, lda: usize,
659/// # b: *const T, ldb: usize, beta: T, c: *mut T, ldc: usize)
660/// # where
661/// # T: Copy + Default + PartialEq + Add<Output = T> + Mul<Output = T>,
662/// # {
663/// # let [m, n, k] = mnk;
664/// # for j in 0..n {
665/// # for i in 0..m {
666/// # let mut acc = T::default();
667/// # for p in 0..k {
668/// # acc = acc + unsafe { *a.add(i + p * lda) * *b.add(p + j * ldb) };
669/// # }
670/// # let cij = unsafe { &mut *c.add(i + j * ldc) };
671/// # *cij = if beta == T::default() { alpha * acc } else { alpha * acc + beta * *cij };
672/// # }
673/// # }
674/// # }
675/// # unsafe extern "C" fn my_dgemm(
676/// # _transa: *const c_char, _transb: *const c_char,
677/// # m: *const c_int, n: *const c_int, k: *const c_int,
678/// # alpha: *const f64, a: *const f64, lda: *const c_int,
679/// # b: *const f64, ldb: *const c_int,
680/// # beta: *const f64, c: *mut f64, ldc: *const c_int,
681/// # ) {
682/// # unsafe {
683/// # let mnk = [*m as usize, *n as usize, *k as usize];
684/// # gemm_nn(mnk, *alpha, a, *lda as usize, b, *ldb as usize, *beta, c, *ldc as usize);
685/// # }
686/// # }
687/// # unsafe extern "C" fn my_zgemm(
688/// # _transa: *const c_char, _transb: *const c_char,
689/// # m: *const c_int, n: *const c_int, k: *const c_int,
690/// # alpha: *const Complex<f64>, a: *const Complex<f64>, lda: *const c_int,
691/// # b: *const Complex<f64>, ldb: *const c_int,
692/// # beta: *const Complex<f64>, c: *mut Complex<f64>, ldc: *const c_int,
693/// # ) {
694/// # unsafe {
695/// # let mnk = [*m as usize, *n as usize, *k as usize];
696/// # gemm_nn(mnk, *alpha, a, *lda as usize, b, *ldb as usize, *beta, c, *ldc as usize);
697/// # }
698/// # }
699///
700/// // `my_dgemm` and `my_zgemm` are the LP64 Fortran BLAS `dgemm_` and `zgemm_`
701/// // to use, e.g. from OpenBLAS. (This example defines naive stand-ins in
702/// // hidden lines so that it runs without a BLAS library.)
703/// unsafe {
704/// set_blas_backend(my_dgemm, my_zgemm);
705/// }
706/// assert_eq!(get_backend_info(), ("External BLAS (LP64)", true, false));
707///
708/// // Calls without an explicit backend handle now go through the injected BLAS
709/// // Column-major data: [[1, 2], [3, 4]] and [[5, 6], [7, 8]]
710/// let a = Matrix::<f64>::from_vec_col_major([2, 2], vec![1.0, 3.0, 2.0, 4.0]).unwrap();
711/// let b = Matrix::<f64>::from_vec_col_major([2, 2], vec![5.0, 7.0, 6.0, 8.0]).unwrap();
712/// let c = matmul(None, &a, &b).unwrap();
713/// assert_eq!(c.host_data().unwrap(), &[19.0, 43.0, 22.0, 50.0]); // [[19, 22], [43, 50]]
714///
715/// // Complex matrices use `zgemm`: (i A) B = i (A B)
716/// let to_complex = |m: &Matrix<f64>, imaginary: bool| {
717/// let data = m.host_data().unwrap().iter().map(|&x| {
718/// if imaginary { Complex::new(0.0, x) } else { Complex::new(x, 0.0) }
719/// });
720/// Matrix::from_vec_col_major([2, 2], data.collect()).unwrap()
721/// };
722/// let ic = matmul(None, &to_complex(&a, true), &to_complex(&b, false)).unwrap();
723/// assert_eq!(ic.host_data().unwrap(), to_complex(&c, true).host_data().unwrap());
724///
725/// clear_blas_backend(); // back to the pure-Rust Faer backend
726/// assert!(!get_backend_info().1);
727/// ```
728pub unsafe fn set_blas_backend(dgemm: DgemmFnPtr, zgemm: ZgemmFnPtr) {
729 let mut dispatcher = BLAS_DISPATCHER.write().unwrap_or_else(|e| e.into_inner());
730 *dispatcher = Box::new(ExternalBlasBackend::new(dgemm, zgemm));
731}
732
733/// Set the process-wide default to an ILP64 BLAS.
734///
735/// # Safety
736/// As for [`set_blas_backend`], with 64-bit integer arguments.
737///
738/// # Example
739/// ```
740/// use sparse_ir::Matrix;
741/// use sparse_ir::gemm::{clear_blas_backend, get_backend_info, matmul, set_ilp64_backend};
742/// # use num_complex::Complex;
743/// # use std::ffi::c_char;
744/// # use std::ops::{Add, Mul};
745/// #
746/// # // Naive column-major C = alpha * A * B + beta * C for transa = transb = 'N'
747/// # // (what sparse-ir passes), standing in for a BLAS library.
748/// # unsafe fn gemm_nn<T>(mnk: [usize; 3], alpha: T, a: *const T, lda: usize,
749/// # b: *const T, ldb: usize, beta: T, c: *mut T, ldc: usize)
750/// # where
751/// # T: Copy + Default + PartialEq + Add<Output = T> + Mul<Output = T>,
752/// # {
753/// # let [m, n, k] = mnk;
754/// # for j in 0..n {
755/// # for i in 0..m {
756/// # let mut acc = T::default();
757/// # for p in 0..k {
758/// # acc = acc + unsafe { *a.add(i + p * lda) * *b.add(p + j * ldb) };
759/// # }
760/// # let cij = unsafe { &mut *c.add(i + j * ldc) };
761/// # *cij = if beta == T::default() { alpha * acc } else { alpha * acc + beta * *cij };
762/// # }
763/// # }
764/// # }
765/// # unsafe extern "C" fn my_dgemm64(
766/// # _transa: *const c_char, _transb: *const c_char,
767/// # m: *const i64, n: *const i64, k: *const i64,
768/// # alpha: *const f64, a: *const f64, lda: *const i64,
769/// # b: *const f64, ldb: *const i64,
770/// # beta: *const f64, c: *mut f64, ldc: *const i64,
771/// # ) {
772/// # unsafe {
773/// # let mnk = [*m as usize, *n as usize, *k as usize];
774/// # gemm_nn(mnk, *alpha, a, *lda as usize, b, *ldb as usize, *beta, c, *ldc as usize);
775/// # }
776/// # }
777/// # unsafe extern "C" fn my_zgemm64(
778/// # _transa: *const c_char, _transb: *const c_char,
779/// # m: *const i64, n: *const i64, k: *const i64,
780/// # alpha: *const Complex<f64>, a: *const Complex<f64>, lda: *const i64,
781/// # b: *const Complex<f64>, ldb: *const i64,
782/// # beta: *const Complex<f64>, c: *mut Complex<f64>, ldc: *const i64,
783/// # ) {
784/// # unsafe {
785/// # let mnk = [*m as usize, *n as usize, *k as usize];
786/// # gemm_nn(mnk, *alpha, a, *lda as usize, b, *ldb as usize, *beta, c, *ldc as usize);
787/// # }
788/// # }
789///
790/// // `my_dgemm64` and `my_zgemm64` are the ILP64 (64-bit integer) Fortran BLAS
791/// // `dgemm_` and `zgemm_` to use. (This example defines naive stand-ins in
792/// // hidden lines so that it runs without an ILP64 BLAS library.)
793/// unsafe {
794/// set_ilp64_backend(my_dgemm64, my_zgemm64);
795/// }
796/// assert_eq!(get_backend_info(), ("External BLAS (ILP64)", true, true));
797///
798/// // Calls without an explicit backend handle now go through the injected BLAS
799/// // Column-major data: [[1, 2], [3, 4]] and [[5, 6], [7, 8]]
800/// let a = Matrix::<f64>::from_vec_col_major([2, 2], vec![1.0, 3.0, 2.0, 4.0]).unwrap();
801/// let b = Matrix::<f64>::from_vec_col_major([2, 2], vec![5.0, 7.0, 6.0, 8.0]).unwrap();
802/// let c = matmul(None, &a, &b).unwrap();
803/// assert_eq!(c.host_data().unwrap(), &[19.0, 43.0, 22.0, 50.0]); // [[19, 22], [43, 50]]
804///
805/// clear_blas_backend(); // back to the pure-Rust Faer backend
806/// let (_, is_external, is_ilp64) = get_backend_info();
807/// assert!(!is_external && !is_ilp64);
808/// ```
809pub unsafe fn set_ilp64_backend(dgemm64: Dgemm64FnPtr, zgemm64: Zgemm64FnPtr) {
810 let mut dispatcher = BLAS_DISPATCHER.write().unwrap_or_else(|e| e.into_inner());
811 *dispatcher = Box::new(ExternalBlas64Backend::new(dgemm64, zgemm64));
812}
813
814/// Reset the process-wide default to the pure Rust faer backend.
815pub fn clear_blas_backend() {
816 let mut dispatcher = BLAS_DISPATCHER.write().unwrap_or_else(|e| e.into_inner());
817 *dispatcher = Box::new(FaerBackend);
818}
819
820/// Returns `(name, is_external, is_ilp64)` of the process-wide default.
821pub fn get_backend_info() -> (&'static str, bool, bool) {
822 let dispatcher = BLAS_DISPATCHER.read().unwrap_or_else(|e| e.into_inner());
823 let name = dispatcher.name();
824 (name, !name.contains("Faer"), dispatcher.is_ilp64())
825}
826
827/// Run `f` with the selected backend (explicit handle or process default).
828fn with_backend<R>(
829 backend: Option<&GemmBackendHandle>,
830 f: impl FnOnce(&dyn GemmBackend) -> R,
831) -> R {
832 match backend {
833 Some(handle) => f(handle.as_ref()),
834 None => {
835 let dispatcher = BLAS_DISPATCHER.read().unwrap_or_else(|e| e.into_inner());
836 f(dispatcher.as_ref())
837 }
838 }
839}
840
841//==============================================================================
842// Scalar dispatch and safe entry points
843//==============================================================================
844
845mod sealed {
846 pub trait Sealed {}
847 impl Sealed for f64 {}
848 impl Sealed for num_complex::Complex<f64> {}
849}
850
851/// Scalars supported by [`gemm`]: `f64` and `Complex<f64>`.
852pub trait GemmScalar:
853 sealed::Sealed
854 + tenferro_tensor::TensorScalar
855 + Copy
856 + Send
857 + Sync
858 + PartialEq
859 + Zero
860 + One
861 + std::ops::Mul<Output = Self>
862 + 'static
863{
864 /// Dispatch to the backend routine for this scalar.
865 ///
866 /// # Safety
867 /// Same contract as [`GemmBackend::dgemm`].
868 #[allow(clippy::too_many_arguments)]
869 unsafe fn gemm_raw(
870 backend: &dyn GemmBackend,
871 transa: Transpose,
872 transb: Transpose,
873 m: usize,
874 n: usize,
875 k: usize,
876 alpha: Self,
877 a: *const Self,
878 lda: usize,
879 b: *const Self,
880 ldb: usize,
881 beta: Self,
882 c: *mut Self,
883 ldc: usize,
884 ) -> Result<(), GemmError>;
885}
886
887impl GemmScalar for f64 {
888 unsafe fn gemm_raw(
889 backend: &dyn GemmBackend,
890 transa: Transpose,
891 transb: Transpose,
892 m: usize,
893 n: usize,
894 k: usize,
895 alpha: f64,
896 a: *const f64,
897 lda: usize,
898 b: *const f64,
899 ldb: usize,
900 beta: f64,
901 c: *mut f64,
902 ldc: usize,
903 ) -> Result<(), GemmError> {
904 // SAFETY: forwarded caller contract.
905 unsafe { backend.dgemm(transa, transb, m, n, k, alpha, a, lda, b, ldb, beta, c, ldc) }
906 }
907}
908
909impl GemmScalar for Complex<f64> {
910 unsafe fn gemm_raw(
911 backend: &dyn GemmBackend,
912 transa: Transpose,
913 transb: Transpose,
914 m: usize,
915 n: usize,
916 k: usize,
917 alpha: Self,
918 a: *const Self,
919 lda: usize,
920 b: *const Self,
921 ldb: usize,
922 beta: Self,
923 c: *mut Self,
924 ldc: usize,
925 ) -> Result<(), GemmError> {
926 // SAFETY: forwarded caller contract.
927 unsafe { backend.zgemm(transa, transb, m, n, k, alpha, a, lda, b, ldb, beta, c, ldc) }
928 }
929}
930
931/// Validate one column-major operand: `rows x cols` stored with leading
932/// dimension `ld` in a buffer of `len` elements.
933fn check_operand(
934 name: &str,
935 rows: usize,
936 cols: usize,
937 ld: usize,
938 len: usize,
939) -> Result<(), GemmError> {
940 if ld < rows.max(1) {
941 return Err(GemmError::InvalidArgument(format!(
942 "leading dimension of {name} is {ld}, expected at least {}",
943 rows.max(1)
944 )));
945 }
946 if rows == 0 || cols == 0 {
947 return Ok(());
948 }
949 let required = ld
950 .checked_mul(cols - 1)
951 .and_then(|x| x.checked_add(rows))
952 .ok_or_else(|| GemmError::InvalidArgument(format!("{name} extent overflows usize")))?;
953 if len < required {
954 return Err(GemmError::InvalidArgument(format!(
955 "{name} buffer has {len} elements, {rows}x{cols} with ld={ld} needs {required}"
956 )));
957 }
958 Ok(())
959}
960
961/// Safe column-major GEMM: `C <- alpha * op(A) * op(B) + beta * C`.
962///
963/// `op(A)` is `m x k`, `op(B)` is `k x n`, and `C` is `m x n`; the stored
964/// operands are addressed through their leading dimensions.
965///
966/// # Errors
967/// Returns [`GemmError::InvalidArgument`] when a buffer is too short or a
968/// leading dimension is too small, and [`GemmError::DimensionOverflow`] when
969/// an LP64 backend cannot represent a dimension.
970#[allow(clippy::too_many_arguments)]
971pub fn gemm<T: GemmScalar>(
972 backend: Option<&GemmBackendHandle>,
973 transa: Transpose,
974 transb: Transpose,
975 m: usize,
976 n: usize,
977 k: usize,
978 alpha: T,
979 a: &[T],
980 lda: usize,
981 b: &[T],
982 ldb: usize,
983 beta: T,
984 c: &mut [T],
985 ldc: usize,
986) -> Result<(), GemmError> {
987 let (ar, ac) = if transa == Transpose::N {
988 (m, k)
989 } else {
990 (k, m)
991 };
992 let (br, bc) = if transb == Transpose::N {
993 (k, n)
994 } else {
995 (n, k)
996 };
997 check_operand("A", ar, ac, lda, a.len())?;
998 check_operand("B", br, bc, ldb, b.len())?;
999 check_operand("C", m, n, ldc, c.len())?;
1000 // SAFETY: extents validated above; `c` is a unique borrow so it cannot
1001 // alias `a` or `b`.
1002 with_backend(backend, |be| unsafe {
1003 T::gemm_raw(
1004 be,
1005 transa,
1006 transb,
1007 m,
1008 n,
1009 k,
1010 alpha,
1011 a.as_ptr(),
1012 lda,
1013 b.as_ptr(),
1014 ldb,
1015 beta,
1016 c.as_mut_ptr(),
1017 ldc,
1018 )
1019 })
1020}
1021
1022/// Below this `pre` extent, a middle-axis contraction is packed into one
1023/// GEMM instead of looping over `post` slices.
1024const SLICE_LOOP_MIN_PRE: usize = 32;
1025
1026/// Matrix-vector products up to this size bypass the default backend.
1027const SMALL_MATVEC_MAX_ELEMS: usize = 128 * 128;
1028const SMALL_MATVEC_MAX_POST: usize = 4;
1029
1030/// Apply a contiguous column-major `m x n` matrix `a` along the middle axis
1031/// of a column-major `[pre, n, post]` array `x`, writing `[pre, m, post]`
1032/// into `y` (overwritten).
1033///
1034/// # Errors
1035/// Returns [`GemmError::InvalidArgument`] when a buffer length does not match
1036/// its extents, or a backend error.
1037#[allow(clippy::too_many_arguments)]
1038pub(crate) fn apply_along_axis<T: GemmScalar>(
1039 backend: Option<&GemmBackendHandle>,
1040 a: &[T],
1041 m: usize,
1042 n: usize,
1043 x: &[T],
1044 pre: usize,
1045 post: usize,
1046 y: &mut [T],
1047) -> Result<(), GemmError> {
1048 let len_ok = |len: usize, rows: usize| {
1049 pre.checked_mul(rows)
1050 .and_then(|v| v.checked_mul(post))
1051 .is_some_and(|v| v == len)
1052 };
1053 if a.len() != m * n || !len_ok(x.len(), n) || !len_ok(y.len(), m) {
1054 return Err(GemmError::InvalidArgument(format!(
1055 "apply_along_axis: A {}, X {}, Y {} elements for m={m}, n={n}, pre={pre}, post={post}",
1056 a.len(),
1057 x.len(),
1058 y.len()
1059 )));
1060 }
1061 if y.is_empty() {
1062 return Ok(());
1063 }
1064 let (one, zero) = (T::one(), T::zero());
1065 let lda = m.max(1);
1066 if pre == 1
1067 && post <= SMALL_MATVEC_MAX_POST
1068 && m * n <= SMALL_MATVEC_MAX_ELEMS
1069 && backend.is_none()
1070 {
1071 // Small matrix-vector products: the default backend's dispatch
1072 // overhead dominates, so accumulate columns of A directly.
1073 for q in 0..post {
1074 let yq = &mut y[q * m..(q + 1) * m];
1075 yq.fill(zero);
1076 for (j, col) in a.chunks_exact(m).enumerate() {
1077 let xj = x[q * n + j];
1078 for (yi, &aij) in yq.iter_mut().zip(col) {
1079 *yi = *yi + aij * xj;
1080 }
1081 }
1082 }
1083 return Ok(());
1084 }
1085 if pre == 1 {
1086 // Y(m x post) = A(m x n) X(n x post)
1087 return gemm(
1088 backend,
1089 Transpose::N,
1090 Transpose::N,
1091 m,
1092 post,
1093 n,
1094 one,
1095 a,
1096 lda,
1097 x,
1098 n.max(1),
1099 zero,
1100 y,
1101 m.max(1),
1102 );
1103 }
1104 if post == 1 || pre >= SLICE_LOOP_MIN_PRE {
1105 // Y_q(pre x m) = X_q(pre x n) A^T for each post index q.
1106 let (xs, ys) = (pre * n, pre * m);
1107 for q in 0..post {
1108 gemm(
1109 backend,
1110 Transpose::N,
1111 Transpose::T,
1112 pre,
1113 m,
1114 n,
1115 one,
1116 &x[q * xs..(q + 1) * xs],
1117 pre,
1118 a,
1119 lda,
1120 zero,
1121 &mut y[q * ys..(q + 1) * ys],
1122 pre,
1123 )?;
1124 }
1125 return Ok(());
1126 }
1127 // Small pre, many post slices: pack X into [n, pre*post], one GEMM, unpack.
1128 let cols = pre * post;
1129 let mut xp = Vec::with_capacity(n * cols);
1130 for q in 0..post {
1131 for p in 0..pre {
1132 let base = p + pre * n * q;
1133 xp.extend((0..n).map(|j| x[base + pre * j]));
1134 }
1135 }
1136 let mut yp = vec![zero; m * cols];
1137 gemm(
1138 backend,
1139 Transpose::N,
1140 Transpose::N,
1141 m,
1142 cols,
1143 n,
1144 one,
1145 a,
1146 lda,
1147 &xp,
1148 n.max(1),
1149 zero,
1150 &mut yp,
1151 m.max(1),
1152 )?;
1153 for q in 0..post {
1154 for p in 0..pre {
1155 let src = &yp[m * (p + pre * q)..m * (p + pre * q + 1)];
1156 let base = p + pre * m * q;
1157 for (i, &v) in src.iter().enumerate() {
1158 y[base + pre * i] = v;
1159 }
1160 }
1161 }
1162 Ok(())
1163}
1164
1165/// Matrix product `A * B` of two contiguous column-major matrices.
1166///
1167/// # Errors
1168/// Returns [`crate::Error::ShapeMismatch`] when the inner dimensions differ,
1169/// [`crate::Error::Tensor`] when an operand is not host-resident compact
1170/// column-major storage, or a backend error.
1171///
1172/// # Example
1173/// ```
1174/// use sparse_ir::Matrix;
1175/// use sparse_ir::gemm::{GemmBackendHandle, matmul};
1176///
1177/// // Column-major data: [[1, 2], [3, 4]] and [[5, 6], [7, 8]]
1178/// let a = Matrix::<f64>::from_vec_col_major([2, 2], vec![1.0, 3.0, 2.0, 4.0]).unwrap();
1179/// let b = Matrix::<f64>::from_vec_col_major([2, 2], vec![5.0, 7.0, 6.0, 8.0]).unwrap();
1180/// let backend = GemmBackendHandle::default();
1181/// let c = matmul(Some(&backend), &a, &b).unwrap();
1182/// assert_eq!(c.host_data().unwrap(), &[19.0, 43.0, 22.0, 50.0]); // [[19, 22], [43, 50]]
1183///
1184/// // A 2x3 times 3x1 product
1185/// let a = Matrix::<f64>::from_vec_col_major([2, 3], vec![1.0, 4.0, 2.0, 5.0, 3.0, 6.0]).unwrap();
1186/// let b = Matrix::<f64>::from_vec_col_major([3, 1], vec![7.0, 8.0, 9.0]).unwrap();
1187/// assert_eq!(matmul(None, &a, &b).unwrap().host_data().unwrap(), &[50.0, 122.0]);
1188/// ```
1189pub fn matmul<T: GemmScalar>(
1190 backend: Option<&GemmBackendHandle>,
1191 a: &crate::Matrix<T>,
1192 b: &crate::Matrix<T>,
1193) -> crate::Result<crate::Matrix<T>> {
1194 let (av, bv) = (a.host_col_major_view()?, b.host_col_major_view()?);
1195 let [m, k] = *av.shape();
1196 let [k2, n] = *bv.shape();
1197 if k != k2 {
1198 return Err(crate::Error::ShapeMismatch {
1199 which: crate::ArrayRole::Input,
1200 expected: vec![k, n],
1201 actual: vec![k2, n],
1202 });
1203 }
1204 let mut c = vec![T::zero(); m * n];
1205 gemm(
1206 backend,
1207 Transpose::N,
1208 Transpose::N,
1209 m,
1210 n,
1211 k,
1212 T::one(),
1213 av.as_slice(),
1214 m.max(1),
1215 bv.as_slice(),
1216 k.max(1),
1217 T::zero(),
1218 &mut c,
1219 m.max(1),
1220 )?;
1221 Ok(crate::Matrix::from_vec_col_major([m, n], c)?)
1222}
1223
1224#[cfg(test)]
1225#[path = "gemm_tests.rs"]
1226mod tests;