Skip to main content

strided_einsum2/
raw_bgemm.rs

1//! Raw borrowed-layout batched GEMM entry points.
2//!
3//! This module is the prepared-replay boundary for callers that already own
4//! validated layout metadata. It keeps the public API independent of a concrete
5//! GEMM backend while still allowing backend modules to provide specialized raw
6//! implementations.
7
8use crate::backend::Backend;
9use crate::{contiguous, Scalar, ScalarBase};
10use strided_view::{Conj, ElementOp, ElementOpApply, RawStridedMut, RawStridedRef};
11
12#[derive(Clone, Copy)]
13pub(crate) struct BgemmGroupLayout {
14    pub(crate) a_sum_end: usize,
15    pub(crate) a_rank: usize,
16    pub(crate) b_ro_end: usize,
17    pub(crate) b_rank: usize,
18    pub(crate) c_ro_end: usize,
19    pub(crate) c_rank: usize,
20    pub(crate) label_len: usize,
21}
22
23pub(crate) fn checked_bgemm_group_layout(
24    n_batch: usize,
25    n_lo: usize,
26    n_ro: usize,
27    n_sum: usize,
28) -> crate::Result<BgemmGroupLayout> {
29    let a_sum_end = n_lo
30        .checked_add(n_sum)
31        .ok_or(strided_view::StridedError::OffsetOverflow)?;
32    let a_rank = a_sum_end
33        .checked_add(n_batch)
34        .ok_or(strided_view::StridedError::OffsetOverflow)?;
35    let b_ro_end = n_sum
36        .checked_add(n_ro)
37        .ok_or(strided_view::StridedError::OffsetOverflow)?;
38    let b_rank = b_ro_end
39        .checked_add(n_batch)
40        .ok_or(strided_view::StridedError::OffsetOverflow)?;
41    let c_ro_end = n_lo
42        .checked_add(n_ro)
43        .ok_or(strided_view::StridedError::OffsetOverflow)?;
44    let c_rank = c_ro_end
45        .checked_add(n_batch)
46        .ok_or(strided_view::StridedError::OffsetOverflow)?;
47    let label_len = c_rank
48        .checked_add(n_sum)
49        .ok_or(strided_view::StridedError::OffsetOverflow)?;
50    Ok(BgemmGroupLayout {
51        a_sum_end,
52        a_rank,
53        b_ro_end,
54        b_rank,
55        c_ro_end,
56        c_rank,
57        label_len,
58    })
59}
60
61/// Batched strided GEMM on raw borrowed layout metadata using the active backend.
62///
63/// This is the raw-layout counterpart to backend-specific `bgemm_strided_into`
64/// functions. It avoids constructing owned-metadata `StridedView` wrappers when
65/// a caller already has borrowed `dims`/`strides`/`offset` descriptors.
66#[allow(clippy::too_many_arguments)]
67pub fn bgemm_raw_strided_into<T>(
68    c: RawStridedMut<'_, T>,
69    a: RawStridedRef<'_, T>,
70    b: RawStridedRef<'_, T>,
71    n_batch: usize,
72    n_lo: usize,
73    n_ro: usize,
74    n_sum: usize,
75    alpha: T,
76    beta: T,
77    conj_a: bool,
78    conj_b: bool,
79) -> crate::Result<()>
80where
81    T: Scalar,
82    crate::backend::ActiveBackend: Backend<T>,
83{
84    validate_bgemm_shapes(&c, &a, &b, n_batch, n_lo, n_ro, n_sum)?;
85    unsafe {
86        bgemm_raw_strided_into_unchecked(
87            c, a, b, n_batch, n_lo, n_ro, n_sum, alpha, beta, conj_a, conj_b,
88        )
89    }
90}
91
92/// Batched strided GEMM on raw borrowed layout metadata without validation.
93///
94/// # Safety
95/// The caller must ensure:
96/// - all raw strided operands are in bounds,
97/// - `n_lo`, `n_ro`, `n_sum`, and `n_batch` partition operand ranks as
98///   `[lo, sum, batch]`, `[sum, ro, batch]`, and `[lo, ro, batch]`,
99/// - matching dimension groups have identical extents,
100/// - `c` does not alias `a` or `b` in a way that violates mutable access.
101#[allow(clippy::too_many_arguments)]
102pub unsafe fn bgemm_raw_strided_into_unchecked<T>(
103    c: RawStridedMut<'_, T>,
104    a: RawStridedRef<'_, T>,
105    b: RawStridedRef<'_, T>,
106    n_batch: usize,
107    n_lo: usize,
108    n_ro: usize,
109    n_sum: usize,
110    alpha: T,
111    beta: T,
112    conj_a: bool,
113    conj_b: bool,
114) -> crate::Result<()>
115where
116    T: Scalar,
117    crate::backend::ActiveBackend: Backend<T>,
118{
119    bgemm_raw_with_backend_into_unchecked::<T, crate::backend::ActiveBackend>(
120        c, a, b, n_batch, n_lo, n_ro, n_sum, alpha, beta, conj_a, conj_b,
121    )
122}
123
124/// Batched strided GEMM on raw borrowed metadata using an explicit backend.
125///
126/// Backend implementations that do not provide a specialized raw path use the
127/// same preparation pipeline as `einsum2_dispatch`: materialize/copy only when
128/// the backend requires it, call `Backend::bgemm_contiguous_into`, then finalize
129/// the destination.
130#[allow(clippy::too_many_arguments)]
131pub fn bgemm_raw_with_backend_into<T, B>(
132    c: RawStridedMut<'_, T>,
133    a: RawStridedRef<'_, T>,
134    b: RawStridedRef<'_, T>,
135    n_batch: usize,
136    n_lo: usize,
137    n_ro: usize,
138    n_sum: usize,
139    alpha: T,
140    beta: T,
141    conj_a: bool,
142    conj_b: bool,
143) -> crate::Result<()>
144where
145    T: ScalarBase + ElementOpApply,
146    B: Backend<T>,
147{
148    validate_bgemm_shapes(&c, &a, &b, n_batch, n_lo, n_ro, n_sum)?;
149    unsafe {
150        bgemm_raw_with_backend_into_unchecked::<T, B>(
151            c, a, b, n_batch, n_lo, n_ro, n_sum, alpha, beta, conj_a, conj_b,
152        )
153    }
154}
155
156/// Unchecked variant of [`bgemm_raw_with_backend_into`].
157///
158/// # Safety
159/// The caller must uphold the same layout and aliasing invariants as
160/// [`bgemm_raw_strided_into_unchecked`].
161#[allow(clippy::too_many_arguments)]
162pub unsafe fn bgemm_raw_with_backend_into_unchecked<T, B>(
163    mut c: RawStridedMut<'_, T>,
164    a: RawStridedRef<'_, T>,
165    b: RawStridedRef<'_, T>,
166    _n_batch: usize,
167    n_lo: usize,
168    n_ro: usize,
169    n_sum: usize,
170    alpha: T,
171    beta: T,
172    conj_a: bool,
173    conj_b: bool,
174) -> crate::Result<()>
175where
176    T: ScalarBase + ElementOpApply,
177    B: Backend<T>,
178{
179    let a_dims = a.dims();
180    let b_dims = b.dims();
181    let lo_dims = &a_dims[..n_lo];
182    let sum_dims = &a_dims[n_lo..n_lo + n_sum];
183    let batch_dims = &a_dims[n_lo + n_sum..];
184    let ro_dims = &b_dims[n_sum..n_sum + n_ro];
185
186    if c.dims().iter().any(|&dim| dim == 0) {
187        return Ok(());
188    }
189    if sum_dims.iter().any(|&dim| dim == 0) {
190        scale_or_zero_raw_mut(&mut c, beta);
191        return Ok(());
192    }
193
194    let use_pool = true;
195    let materialize = if B::MATERIALIZES_CONJ {
196        Some(Conj::apply as fn(T) -> T)
197    } else {
198        None
199    };
200
201    let a_op = contiguous::prepare_input_raw(
202        &a,
203        n_lo,
204        n_sum,
205        conj_a,
206        B::REQUIRES_UNIT_STRIDE,
207        use_pool,
208        materialize,
209    )?;
210    let b_op = contiguous::prepare_input_raw(
211        &b,
212        n_sum,
213        n_ro,
214        conj_b,
215        B::REQUIRES_UNIT_STRIDE,
216        use_pool,
217        materialize,
218    )?;
219    let mut c_op = contiguous::prepare_output_raw(
220        &mut c,
221        n_lo,
222        n_ro,
223        beta,
224        B::REQUIRES_UNIT_STRIDE,
225        use_pool,
226    )?;
227
228    let m: usize = lo_dims.iter().product::<usize>().max(1);
229    let k: usize = sum_dims.iter().product::<usize>().max(1);
230    let n: usize = ro_dims.iter().product::<usize>().max(1);
231
232    B::bgemm_contiguous_into(&mut c_op, &a_op, &b_op, batch_dims, m, n, k, alpha, beta)?;
233    c_op.finalize_raw_into(&mut c)?;
234
235    Ok(())
236}
237
238pub(crate) fn validate_bgemm_shapes<T, U>(
239    c: &RawStridedMut<'_, U>,
240    a: &RawStridedRef<'_, T>,
241    b: &RawStridedRef<'_, T>,
242    n_batch: usize,
243    n_lo: usize,
244    n_ro: usize,
245    n_sum: usize,
246) -> crate::Result<()> {
247    let groups = checked_bgemm_group_layout(n_batch, n_lo, n_ro, n_sum)?;
248    if a.dims().len() != groups.a_rank {
249        return Err(strided_view::StridedError::RankMismatch(groups.a_rank, a.dims().len()).into());
250    }
251    if b.dims().len() != groups.b_rank {
252        return Err(strided_view::StridedError::RankMismatch(groups.b_rank, b.dims().len()).into());
253    }
254    if c.dims().len() != groups.c_rank {
255        return Err(strided_view::StridedError::RankMismatch(groups.c_rank, c.dims().len()).into());
256    }
257
258    let lo_dims = &a.dims()[..n_lo];
259    let sum_dims = &a.dims()[n_lo..groups.a_sum_end];
260    let batch_dims = &a.dims()[groups.a_sum_end..];
261    let ro_dims = &b.dims()[n_sum..groups.b_ro_end];
262
263    if &b.dims()[..n_sum] != sum_dims {
264        return Err(strided_view::StridedError::ShapeMismatch(
265            sum_dims.to_vec(),
266            b.dims()[..n_sum].to_vec(),
267        )
268        .into());
269    }
270    if &b.dims()[groups.b_ro_end..] != batch_dims {
271        return Err(strided_view::StridedError::ShapeMismatch(
272            batch_dims.to_vec(),
273            b.dims()[groups.b_ro_end..].to_vec(),
274        )
275        .into());
276    }
277    if &c.dims()[..n_lo] != lo_dims {
278        return Err(strided_view::StridedError::ShapeMismatch(
279            lo_dims.to_vec(),
280            c.dims()[..n_lo].to_vec(),
281        )
282        .into());
283    }
284    if &c.dims()[n_lo..groups.c_ro_end] != ro_dims {
285        return Err(strided_view::StridedError::ShapeMismatch(
286            ro_dims.to_vec(),
287            c.dims()[n_lo..groups.c_ro_end].to_vec(),
288        )
289        .into());
290    }
291    if &c.dims()[groups.c_ro_end..] != batch_dims {
292        return Err(strided_view::StridedError::ShapeMismatch(
293            batch_dims.to_vec(),
294            c.dims()[groups.c_ro_end..].to_vec(),
295        )
296        .into());
297    }
298    Ok(())
299}
300
301pub(crate) fn scale_or_zero_raw_mut<T: ScalarBase>(c: &mut RawStridedMut<'_, T>, beta: T) {
302    if c.dims().iter().any(|&dim| dim == 0) {
303        return;
304    }
305
306    fn visit<T: ScalarBase>(
307        ptr: *mut T,
308        dims: &[usize],
309        strides: &[isize],
310        axis: usize,
311        offset: isize,
312        beta: T,
313        zero: T,
314    ) {
315        if axis == dims.len() {
316            unsafe {
317                let dst = ptr.offset(offset);
318                if beta == zero {
319                    *dst = zero;
320                } else {
321                    *dst = beta * *dst;
322                }
323            }
324            return;
325        }
326
327        for i in 0..dims[axis] {
328            visit(
329                ptr,
330                dims,
331                strides,
332                axis + 1,
333                offset + i as isize * strides[axis],
334                beta,
335                zero,
336            );
337        }
338    }
339
340    visit(c.as_mut_ptr(), c.dims(), c.strides(), 0, 0, beta, T::zero());
341}
342
343#[cfg(test)]
344mod tests {
345    use super::*;
346
347    fn raw_bgemm_2x2<T>(one: T, zero: T) -> Vec<T>
348    where
349        T: Scalar,
350        crate::backend::ActiveBackend: Backend<T>,
351        T: From<f32>,
352    {
353        let dims = [2, 2];
354        let strides = [2, 1];
355        let a_data = [T::from(1.0), T::from(2.0), T::from(3.0), T::from(4.0)];
356        let b_data = [T::from(5.0), T::from(6.0), T::from(7.0), T::from(8.0)];
357        let mut c_data = vec![zero; 4];
358        let a = RawStridedRef::new(&a_data, &dims, &strides, 0).unwrap();
359        let b = RawStridedRef::new(&b_data, &dims, &strides, 0).unwrap();
360        let c = RawStridedMut::new(&mut c_data, &dims, &strides, 0).unwrap();
361        bgemm_raw_strided_into(c, a, b, 0, 1, 1, 1, one, zero, false, false).unwrap();
362        c_data
363    }
364
365    #[test]
366    fn raw_bgemm_active_backend_f64() {
367        assert_eq!(raw_bgemm_2x2(1.0f64, 0.0), vec![19.0, 22.0, 43.0, 50.0]);
368    }
369
370    #[test]
371    fn raw_bgemm_active_backend_f32() {
372        assert_eq!(raw_bgemm_2x2(1.0f32, 0.0), vec![19.0f32, 22.0, 43.0, 50.0]);
373    }
374
375    #[test]
376    fn raw_bgemm_active_backend_complex_conj() {
377        use num_complex::Complex64;
378
379        let i = Complex64::i();
380        let dims = [2, 2];
381        let strides = [2, 1];
382        let a_data = [
383            Complex64::new(1.0, 0.0) + i,
384            Complex64::new(2.0, 0.0),
385            Complex64::new(3.0, 0.0),
386            Complex64::new(4.0, 0.0) - i,
387        ];
388        let b_data = [
389            Complex64::new(1.0, 0.0),
390            Complex64::new(0.0, 0.0),
391            Complex64::new(0.0, 0.0),
392            Complex64::new(1.0, 0.0),
393        ];
394        let mut c_data = vec![Complex64::new(0.0, 0.0); 4];
395        let a = RawStridedRef::new(&a_data, &dims, &strides, 0).unwrap();
396        let b = RawStridedRef::new(&b_data, &dims, &strides, 0).unwrap();
397        let c = RawStridedMut::new(&mut c_data, &dims, &strides, 0).unwrap();
398        bgemm_raw_strided_into(
399            c,
400            a,
401            b,
402            0,
403            1,
404            1,
405            1,
406            Complex64::new(1.0, 0.0),
407            Complex64::new(0.0, 0.0),
408            true,
409            false,
410        )
411        .unwrap();
412        assert_eq!(
413            c_data,
414            vec![
415                Complex64::new(1.0, -1.0),
416                Complex64::new(2.0, 0.0),
417                Complex64::new(3.0, 0.0),
418                Complex64::new(4.0, 1.0),
419            ]
420        );
421    }
422
423    #[test]
424    fn raw_bgemm_active_backend_checked_shape_mismatch() {
425        let a_dims = [2, 2];
426        let b_dims = [3, 2];
427        let c_dims = [2, 2];
428        let a_strides = [2, 1];
429        let b_strides = [2, 1];
430        let c_strides = [2, 1];
431        let a_data = [1.0, 2.0, 3.0, 4.0];
432        let b_data = [0.0; 6];
433        let mut c_data = [0.0; 4];
434        let a = RawStridedRef::new(&a_data, &a_dims, &a_strides, 0).unwrap();
435        let b = RawStridedRef::new(&b_data, &b_dims, &b_strides, 0).unwrap();
436        let c = RawStridedMut::new(&mut c_data, &c_dims, &c_strides, 0).unwrap();
437        let err = bgemm_raw_strided_into(c, a, b, 0, 1, 1, 1, 1.0, 0.0, false, false).unwrap_err();
438        assert!(matches!(
439            err,
440            crate::EinsumError::Strided(strided_view::StridedError::ShapeMismatch(_, _))
441        ));
442    }
443
444    #[test]
445    fn raw_bgemm_explicit_backend_checked_rank_mismatch() {
446        let a_dims = [2, 2];
447        let b_dims = [2, 2];
448        let c_dims = [2];
449        let strides = [2, 1];
450        let c_strides = [1];
451        let a_data = [1.0, 2.0, 3.0, 4.0];
452        let b_data = [5.0, 6.0, 7.0, 8.0];
453        let mut c_data = [0.0; 2];
454        let a = RawStridedRef::new(&a_data, &a_dims, &strides, 0).unwrap();
455        let b = RawStridedRef::new(&b_data, &b_dims, &strides, 0).unwrap();
456        let c = RawStridedMut::new(&mut c_data, &c_dims, &c_strides, 0).unwrap();
457        let err = bgemm_raw_with_backend_into::<f64, crate::backend::ActiveBackend>(
458            c, a, b, 0, 1, 1, 1, 1.0, 0.0, false, false,
459        )
460        .unwrap_err();
461        assert!(matches!(
462            err,
463            crate::EinsumError::Strided(strided_view::StridedError::RankMismatch(2, 1))
464        ));
465    }
466
467    #[test]
468    fn raw_bgemm_zero_sum_scales_destination() {
469        let a_dims = [2, 0];
470        let b_dims = [0, 2];
471        let c_dims = [2, 2];
472        let a_strides = [0, 0];
473        let b_strides = [0, 0];
474        let c_strides = [2, 1];
475        let a_data = [0.0; 1];
476        let b_data = [0.0; 1];
477        let mut c_data = [1.0, 2.0, 3.0, 4.0];
478        let a = RawStridedRef::new(&a_data, &a_dims, &a_strides, 0).unwrap();
479        let b = RawStridedRef::new(&b_data, &b_dims, &b_strides, 0).unwrap();
480        let c = RawStridedMut::new(&mut c_data, &c_dims, &c_strides, 0).unwrap();
481
482        bgemm_raw_strided_into(c, a, b, 0, 1, 1, 1, 1.0, 2.0, false, false).unwrap();
483
484        assert_eq!(c_data, [2.0, 4.0, 6.0, 8.0]);
485    }
486
487    #[test]
488    fn raw_bgemm_zero_sum_beta_zero_clears_destination() {
489        let a_dims = [2, 0];
490        let b_dims = [0, 2];
491        let c_dims = [2, 2];
492        let a_strides = [0, 0];
493        let b_strides = [0, 0];
494        let c_strides = [2, 1];
495        let a_data = [0.0; 1];
496        let b_data = [0.0; 1];
497        let mut c_data = [1.0, 2.0, 3.0, 4.0];
498        let a = RawStridedRef::new(&a_data, &a_dims, &a_strides, 0).unwrap();
499        let b = RawStridedRef::new(&b_data, &b_dims, &b_strides, 0).unwrap();
500        let c = RawStridedMut::new(&mut c_data, &c_dims, &c_strides, 0).unwrap();
501
502        bgemm_raw_strided_into(c, a, b, 0, 1, 1, 1, 1.0, 0.0, false, false).unwrap();
503
504        assert_eq!(c_data, [0.0, 0.0, 0.0, 0.0]);
505    }
506
507    #[test]
508    fn raw_bgemm_empty_output_is_noop() {
509        let a_dims = [0, 2];
510        let b_dims = [2, 2];
511        let c_dims = [0, 2];
512        let a_strides = [2, 1];
513        let b_strides = [2, 1];
514        let c_strides = [2, 1];
515        let a_data = [1.0, 2.0];
516        let b_data = [3.0, 4.0, 5.0, 6.0];
517        let mut c_data = [7.0, 8.0, 9.0, 10.0];
518        let expected = c_data;
519        let a = RawStridedRef::new(&a_data, &a_dims, &a_strides, 0).unwrap();
520        let b = RawStridedRef::new(&b_data, &b_dims, &b_strides, 0).unwrap();
521        let c = RawStridedMut::new(&mut c_data, &c_dims, &c_strides, 0).unwrap();
522
523        bgemm_raw_strided_into(c, a, b, 0, 1, 1, 1, 1.0, 1.0, false, false).unwrap();
524
525        assert_eq!(c_data, expected);
526    }
527
528    #[test]
529    fn raw_bgemm_noncontiguous_output_writes_back() {
530        let a_dims = [2, 2];
531        let b_dims = [2, 2];
532        let c_dims = [2, 2];
533        let a_strides = [2, 1];
534        let b_strides = [2, 1];
535        let c_strides = [1, 3];
536        let a_data = [1.0, 2.0, 3.0, 4.0];
537        let b_data = [5.0, 6.0, 7.0, 8.0];
538        let mut c_data = [0.0; 8];
539        let a = RawStridedRef::new(&a_data, &a_dims, &a_strides, 0).unwrap();
540        let b = RawStridedRef::new(&b_data, &b_dims, &b_strides, 0).unwrap();
541        let c = RawStridedMut::new(&mut c_data, &c_dims, &c_strides, 1).unwrap();
542
543        bgemm_raw_strided_into(c, a, b, 0, 1, 1, 1, 1.0, 0.0, false, false).unwrap();
544
545        assert_eq!(c_data[1], 19.0);
546        assert_eq!(c_data[4], 22.0);
547        assert_eq!(c_data[2], 43.0);
548        assert_eq!(c_data[5], 50.0);
549        assert_eq!(c_data[0], 0.0);
550        assert_eq!(c_data[3], 0.0);
551    }
552}