Skip to main content

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;