Skip to main content

strided_einsum2/
uninit.rs

1//! Overwrite-only contraction entry points.
2//!
3//! This module deliberately does not reuse the initialized `beta` path.  The
4//! destination is borrowed as `MaybeUninit<T>` until the last logical element
5//! has been written, so a provider failure cannot expose a partially
6//! initialized slice as `T`.
7
8use std::collections::HashSet;
9use std::mem::MaybeUninit;
10
11use strided_kernel::ExecContext;
12use strided_view::{ElementOp, RawStridedMut, RawStridedRef, StridedView};
13
14use crate::{AxisId, Einsum2Plan, EinsumError, Result, ScalarBase};
15
16/// Naive overwrite kernel used by the private backend contract and as the
17/// no-provider implementation. It never reads C.
18#[cfg(not(any(feature = "blas", feature = "blas-inject")))]
19pub(crate) fn bgemm_contiguous_naive<T>(
20    c: &mut crate::contiguous::UninitContiguousOperand<'_, '_, T>,
21    a: &crate::contiguous::ContiguousOperand<T>,
22    b: &crate::contiguous::ContiguousOperand<T>,
23    batch_dims: &[usize],
24    m: usize,
25    n: usize,
26    k: usize,
27    alpha: T,
28    _ctx: &ExecContext,
29) -> strided_view::Result<()>
30where
31    T: ScalarBase + strided_view::ElementOpApply,
32{
33    let mut batch = crate::util::MultiIndex::new(batch_dims);
34    while batch.next().is_some() {
35        let a_base = batch.offset(a.batch_strides());
36        let b_base = batch.offset(b.batch_strides());
37        let c_base = batch.offset(c.batch_strides());
38        for i in 0..m {
39            for j in 0..n {
40                let mut acc = T::zero();
41                for l in 0..k {
42                    let mut av = unsafe {
43                        *a.ptr().offset(
44                            a_base + i as isize * a.row_stride() + l as isize * a.col_stride(),
45                        )
46                    };
47                    let mut bv = unsafe {
48                        *b.ptr().offset(
49                            b_base + l as isize * b.row_stride() + j as isize * b.col_stride(),
50                        )
51                    };
52                    if a.conj() {
53                        av = strided_view::Conj::apply(av);
54                    }
55                    if b.conj() {
56                        bv = strided_view::Conj::apply(bv);
57                    }
58                    acc = acc + av * bv;
59                }
60                let offset = c_base + i as isize * c.row_stride() + j as isize * c.col_stride();
61                unsafe {
62                    c.ptr().offset(offset).write(MaybeUninit::new(alpha * acc));
63                }
64            }
65        }
66    }
67    Ok(())
68}
69
70#[cfg(any(feature = "blas", feature = "blas-inject"))]
71fn zero_raw_uninit<T: ScalarBase>(dest: &mut RawStridedMut<'_, MaybeUninit<T>>) -> Result<()> {
72    fn visit<T: ScalarBase>(
73        dest: &mut RawStridedMut<'_, MaybeUninit<T>>,
74        dims: &[usize],
75        strides: &[isize],
76        axis: usize,
77        offset: isize,
78    ) -> Result<()> {
79        if axis == dims.len() {
80            let relative = offset
81                .checked_sub(dest.offset())
82                .ok_or(strided_view::StridedError::OffsetOverflow)?;
83            unsafe {
84                dest.as_mut_ptr()
85                    .offset(relative)
86                    .write(MaybeUninit::new(T::zero()));
87            }
88            return Ok(());
89        }
90        for i in 0..dims[axis] {
91            let next = checked_offset(offset, i, strides[axis])?;
92            visit(dest, dims, strides, axis + 1, next)?;
93        }
94        Ok(())
95    }
96    visit(dest, dest.dims(), dest.strides(), 0, dest.offset())
97}
98
99#[cfg(any(feature = "blas", feature = "blas-inject"))]
100fn bgemm_raw_backend<T, B>(
101    mut dest: &mut RawStridedMut<'_, MaybeUninit<T>>,
102    a: &RawStridedRef<'_, T>,
103    b: &RawStridedRef<'_, T>,
104    _n_batch: usize,
105    n_lo: usize,
106    n_ro: usize,
107    n_sum: usize,
108    alpha: T,
109    ctx: &ExecContext,
110) -> Result<()>
111where
112    T: ScalarBase + strided_view::ElementOpApply,
113    B: crate::backend::Backend<T> + crate::backend::OverwriteBackend<T>,
114{
115    let (groups, m, k, n) = preflight_raw_bgemm(&mut dest, a, b, _n_batch, n_lo, n_ro, n_sum)?;
116    if dest.dims().iter().any(|&d| d == 0) {
117        return Ok(());
118    }
119    let sum_dims = &a.dims()[n_lo..groups.a_sum_end];
120    let batch_dims = &a.dims()[groups.a_sum_end..];
121    if sum_dims.iter().any(|&d| d == 0) {
122        zero_raw_uninit(&mut dest)?;
123        return Ok(());
124    }
125    let a_op = crate::contiguous::prepare_input_raw(
126        a,
127        n_lo,
128        n_sum,
129        false,
130        B::REQUIRES_UNIT_STRIDE,
131        true,
132        None,
133    )?;
134    let b_op = crate::contiguous::prepare_input_raw(
135        b,
136        n_sum,
137        n_ro,
138        false,
139        B::REQUIRES_UNIT_STRIDE,
140        true,
141        None,
142    )?;
143    let mut c_op = crate::contiguous::prepare_output_raw_uninit(
144        &mut dest,
145        n_lo,
146        n_ro,
147        B::REQUIRES_UNIT_STRIDE,
148    )?;
149    B::bgemm_contiguous_overwrite(&mut c_op, &a_op, &b_op, batch_dims, m, n, k, alpha, ctx)?;
150    c_op.finalize()?;
151    Ok(())
152}
153
154fn checked_offset(offset: isize, index: usize, stride: isize) -> Result<isize> {
155    let term = (index as isize)
156        .checked_mul(stride)
157        .ok_or(strided_view::StridedError::OffsetOverflow)?;
158    offset
159        .checked_add(term)
160        .ok_or(strided_view::StridedError::OffsetOverflow)
161        .map_err(Into::into)
162}
163
164fn visit_offsets(
165    dims: &[usize],
166    strides: &[isize],
167    axis: usize,
168    offset: isize,
169    seen: &mut HashSet<isize>,
170) -> Result<()> {
171    if axis == dims.len() {
172        if !seen.insert(offset) {
173            return Err(strided_view::StridedError::NonInjectiveOutputLayout.into());
174        }
175        return Ok(());
176    }
177    for index in 0..dims[axis] {
178        visit_offsets(
179            dims,
180            strides,
181            axis + 1,
182            checked_offset(offset, index, strides[axis])?,
183            seen,
184        )?;
185    }
186    Ok(())
187}
188
189fn validate_output<T>(dest: &mut RawStridedMut<'_, MaybeUninit<T>>) -> Result<()> {
190    let mut seen = HashSet::new();
191    visit_offsets(dest.dims(), dest.strides(), 0, dest.offset(), &mut seen)
192}
193
194fn ranges_overlap<T, U>(a_ptr: *const T, a_len: usize, b_ptr: *const U, b_len: usize) -> bool {
195    let a_start = a_ptr as usize;
196    let b_start = b_ptr as usize;
197    let a_bytes = a_len.saturating_mul(std::mem::size_of::<T>());
198    let b_bytes = b_len.saturating_mul(std::mem::size_of::<U>());
199    let a_end = a_start.saturating_add(a_bytes);
200    let b_end = b_start.saturating_add(b_bytes);
201    a_start < b_end && b_start < a_end
202}
203
204fn validate_no_overlap<T, OpA, OpB>(
205    dest: &mut RawStridedMut<'_, MaybeUninit<T>>,
206    a: &StridedView<'_, T, OpA>,
207    b: &StridedView<'_, T, OpB>,
208) -> Result<()>
209where
210    T: Copy,
211    OpA: ElementOp<T>,
212    OpB: ElementOp<T>,
213{
214    let d = dest.data_mut();
215    if ranges_overlap(d.as_ptr(), d.len(), a.data().as_ptr(), a.data().len())
216        || ranges_overlap(d.as_ptr(), d.len(), b.data().as_ptr(), b.data().len())
217    {
218        return Err(strided_view::StridedError::OverlappingInputOutput { input: 0 }.into());
219    }
220    Ok(())
221}
222
223/// Complete raw GEMM preflight, before labels, temporaries, or provider work.
224fn preflight_raw_bgemm<T: Copy>(
225    dest: &mut RawStridedMut<'_, MaybeUninit<T>>,
226    a: &RawStridedRef<'_, T>,
227    b: &RawStridedRef<'_, T>,
228    n_batch: usize,
229    n_lo: usize,
230    n_ro: usize,
231    n_sum: usize,
232) -> Result<(crate::raw_bgemm::BgemmGroupLayout, usize, usize, usize)> {
233    let groups = crate::raw_bgemm::checked_bgemm_group_layout(n_batch, n_lo, n_ro, n_sum)?;
234    crate::raw_bgemm::validate_bgemm_shapes(dest, a, b, n_batch, n_lo, n_ro, n_sum)?;
235    let av: StridedView<'_, T> =
236        unsafe { StridedView::new_unchecked(a.data(), a.dims(), a.strides(), a.offset()) };
237    let bv: StridedView<'_, T> =
238        unsafe { StridedView::new_unchecked(b.data(), b.dims(), b.strides(), b.offset()) };
239    validate_output(dest)?;
240    validate_no_overlap(dest, &av, &bv)?;
241    let m = a.dims()[..n_lo]
242        .iter()
243        .try_fold(1usize, |v, &d| v.checked_mul(d))
244        .ok_or(strided_view::StridedError::OffsetOverflow)?
245        .max(1);
246    let k = a.dims()[n_lo..groups.a_sum_end]
247        .iter()
248        .try_fold(1usize, |v, &d| v.checked_mul(d))
249        .ok_or(strided_view::StridedError::OffsetOverflow)?
250        .max(1);
251    let n = b.dims()[n_sum..groups.b_ro_end]
252        .iter()
253        .try_fold(1usize, |v, &d| v.checked_mul(d))
254        .ok_or(strided_view::StridedError::OffsetOverflow)?
255        .max(1);
256    #[cfg(any(feature = "blas", feature = "blas-inject"))]
257    for value in [m, k, n] {
258        i32::try_from(value).map_err(|_| strided_view::StridedError::OffsetOverflow)?;
259    }
260    Ok((groups, m, k, n))
261}
262
263fn validate_labels<T, OpA, OpB, ID>(
264    plan: &Einsum2Plan<ID>,
265    dest: &RawStridedMut<'_, MaybeUninit<T>>,
266    a: &StridedView<'_, T, OpA>,
267    b: &StridedView<'_, T, OpB>,
268    ic: &[ID],
269    ia: &[ID],
270    ib: &[ID],
271) -> Result<()>
272where
273    T: Copy,
274    OpA: ElementOp<T>,
275    OpB: ElementOp<T>,
276    ID: AxisId,
277{
278    if ia.len() != a.dims().len() || ib.len() != b.dims().len() || ic.len() != dest.dims().len() {
279        return Err(EinsumError::OutputShapeMismatch {
280            expected: vec![ic.len()],
281            got: vec![dest.dims().len()],
282        });
283    }
284    let dim = |labels: &[ID], dims: &[usize], id: &ID| {
285        labels.iter().position(|x| x == id).map(|i| dims[i])
286    };
287    for (axis, id) in ic.iter().enumerate() {
288        let expected = dim(ia, a.dims(), id).or_else(|| dim(ib, b.dims(), id));
289        if expected != Some(dest.dims()[axis]) {
290            return Err(EinsumError::OutputShapeMismatch {
291                expected: ic
292                    .iter()
293                    .map(|x| {
294                        dim(ia, a.dims(), x)
295                            .or_else(|| dim(ib, b.dims(), x))
296                            .unwrap_or(0)
297                    })
298                    .collect(),
299                got: dest.dims().to_vec(),
300            });
301        }
302    }
303    for id in plan.batch.iter().chain(plan.sum.iter()) {
304        let da = dim(ia, a.dims(), id).ok_or_else(|| {
305            EinsumError::InvalidDotGeneralConfig(format!(
306                "planned axis {:?} is absent from lhs",
307                id
308            ))
309        })?;
310        let db = dim(ib, b.dims(), id).ok_or_else(|| {
311            EinsumError::InvalidDotGeneralConfig(format!(
312                "planned axis {:?} is absent from rhs",
313                id
314            ))
315        })?;
316        if da != db {
317            return Err(EinsumError::DimensionMismatch {
318                axis: format!("{:?}", id),
319                dim_a: da,
320                dim_b: db,
321            });
322        }
323    }
324    Ok(())
325}
326
327#[cfg(all(
328    not(any(feature = "blas", feature = "blas-inject")),
329    not(feature = "faer")
330))]
331fn visit_sum<T, OpA, OpB, ID>(
332    sum_ids: &[ID],
333    axis: usize,
334    a_idx: &mut [usize],
335    b_idx: &mut [usize],
336    ia: &[ID],
337    ib: &[ID],
338    a: &StridedView<'_, T, OpA>,
339    b: &StridedView<'_, T, OpB>,
340    acc: &mut T,
341) where
342    T: ScalarBase,
343    OpA: ElementOp<T>,
344    OpB: ElementOp<T>,
345    ID: AxisId,
346{
347    if axis == sum_ids.len() {
348        *acc = *acc + a.get(a_idx) * b.get(b_idx);
349        return;
350    }
351    let id = &sum_ids[axis];
352    let ai = ia.iter().position(|x| x == id);
353    let bi = ib.iter().position(|x| x == id);
354    let dim = ai
355        .map(|i| a.dims()[i])
356        .or_else(|| bi.map(|i| b.dims()[i]))
357        .unwrap_or(0);
358    for i in 0..dim {
359        if let Some(ai) = ai {
360            a_idx[ai] = i;
361        }
362        if let Some(bi) = bi {
363            b_idx[bi] = i;
364        }
365        visit_sum(sum_ids, axis + 1, a_idx, b_idx, ia, ib, a, b, acc);
366    }
367}
368
369#[cfg(all(
370    not(any(feature = "blas", feature = "blas-inject")),
371    not(feature = "faer")
372))]
373fn visit_output<T, OpA, OpB, ID>(
374    axis: usize,
375    out_idx: &mut [usize],
376    dest: &mut RawStridedMut<'_, MaybeUninit<T>>,
377    a_idx: &mut [usize],
378    b_idx: &mut [usize],
379    ic: &[ID],
380    ia: &[ID],
381    ib: &[ID],
382    reduction_ids: &[ID],
383    a: &StridedView<'_, T, OpA>,
384    b: &StridedView<'_, T, OpB>,
385    alpha: T,
386) -> Result<()>
387where
388    T: ScalarBase,
389    OpA: ElementOp<T>,
390    OpB: ElementOp<T>,
391    ID: AxisId,
392{
393    if axis == out_idx.len() {
394        for (pos, id) in ic.iter().enumerate() {
395            if let Some(ai) = ia.iter().position(|x| x == id) {
396                a_idx[ai] = out_idx[pos];
397            }
398            if let Some(bi) = ib.iter().position(|x| x == id) {
399                b_idx[bi] = out_idx[pos];
400            }
401        }
402        let mut value = T::zero();
403        visit_sum(reduction_ids, 0, a_idx, b_idx, ia, ib, a, b, &mut value);
404        let mut offset = dest.offset();
405        for (&idx, &stride) in out_idx.iter().zip(dest.strides()) {
406            offset = checked_offset(offset, idx, stride)?;
407        }
408        let relative = offset
409            .checked_sub(dest.offset())
410            .ok_or(strided_view::StridedError::OffsetOverflow)?;
411        unsafe {
412            dest.as_mut_ptr()
413                .offset(relative)
414                .write(MaybeUninit::new(alpha * value))
415        };
416        return Ok(());
417    }
418    for i in 0..dest.dims()[axis] {
419        out_idx[axis] = i;
420        visit_output(
421            axis + 1,
422            out_idx,
423            dest,
424            a_idx,
425            b_idx,
426            ic,
427            ia,
428            ib,
429            reduction_ids,
430            a,
431            b,
432            alpha,
433        )?;
434    }
435    Ok(())
436}
437
438/// Compute an einsum into a genuinely uninitialized destination.
439#[allow(clippy::too_many_arguments)]
440#[cfg(not(any(feature = "blas", feature = "blas-inject")))]
441pub fn einsum2_into_uninit<T, OpA, OpB, ID>(
442    dest: &mut RawStridedMut<'_, MaybeUninit<T>>,
443    a: &StridedView<'_, T, OpA>,
444    b: &StridedView<'_, T, OpB>,
445    ic: &[ID],
446    ia: &[ID],
447    ib: &[ID],
448    alpha: T,
449    _ctx: &ExecContext,
450) -> Result<()>
451where
452    T: ScalarBase,
453    OpA: ElementOp<T>,
454    OpB: ElementOp<T>,
455    ID: AxisId,
456{
457    let plan = Einsum2Plan::new(ia, ib, ic)?;
458    validate_labels(&plan, dest, a, b, ic, ia, ib)?;
459    validate_output(dest)?;
460    validate_no_overlap(dest, a, b)?;
461    #[cfg(feature = "faer")]
462    {
463        let _ = alpha;
464        return Err(EinsumError::Unsupported(
465            "Faer does not yet expose a MaybeUninit-safe overwrite GEMM API; see strided-rs#195"
466                .to_owned(),
467        ));
468    }
469    #[cfg(not(feature = "faer"))]
470    {
471        if dest.dims().iter().any(|&d| d == 0) {
472            return Ok(());
473        }
474        let mut out_idx = vec![0; dest.dims().len()];
475        let mut a_idx = vec![0; a.dims().len()];
476        let mut b_idx = vec![0; b.dims().len()];
477        let mut reduction_ids = plan.sum.clone();
478        for id in ia {
479            if !ic.contains(id) && !reduction_ids.contains(id) {
480                reduction_ids.push(id.clone());
481            }
482        }
483        for id in ib {
484            if !ic.contains(id) && !reduction_ids.contains(id) {
485                reduction_ids.push(id.clone());
486            }
487        }
488        visit_output(
489            0,
490            &mut out_idx,
491            dest,
492            &mut a_idx,
493            &mut b_idx,
494            ic,
495            ia,
496            ib,
497            &reduction_ids,
498            a,
499            b,
500            alpha,
501        )?;
502        Ok(())
503    }
504}
505
506/// BLAS-backed overwrite path. All public validation happens before the
507/// canonical descriptors are prepared or a temporary is allocated.
508#[allow(clippy::too_many_arguments)]
509#[cfg(any(feature = "blas", feature = "blas-inject"))]
510pub fn einsum2_into_uninit<T, OpA, OpB, ID>(
511    dest: &mut RawStridedMut<'_, MaybeUninit<T>>,
512    a: &StridedView<'_, T, OpA>,
513    b: &StridedView<'_, T, OpB>,
514    ic: &[ID],
515    ia: &[ID],
516    ib: &[ID],
517    alpha: T,
518    ctx: &ExecContext,
519) -> Result<()>
520where
521    T: crate::Scalar,
522    OpA: ElementOp<T> + 'static,
523    OpB: ElementOp<T> + 'static,
524    ID: AxisId,
525{
526    let plan = Einsum2Plan::new(ia, ib, ic)?;
527    validate_labels(&plan, dest, a, b, ic, ia, ib)?;
528    validate_output(dest)?;
529    validate_no_overlap(dest, a, b)?;
530    if dest.dims().iter().any(|&d| d == 0) {
531        return Ok(());
532    }
533
534    let left_trace = crate::trace::find_trace_indices(ia, ib, ic);
535    let (a_buf, conj_a) = if !left_trace.is_empty() {
536        (
537            Some(crate::trace::reduce_trace_axes(a, &left_trace)?),
538            false,
539        )
540    } else {
541        (None, crate::op_is_conj::<OpA>())
542    };
543    let a_view: StridedView<'_, T> = match a_buf.as_ref() {
544        Some(buf) => buf.view(),
545        None => StridedView::new(a.data(), a.dims(), a.strides(), a.offset())?,
546    };
547    let right_trace = crate::trace::find_trace_indices(ib, ia, ic);
548    let (b_buf, conj_b) = if !right_trace.is_empty() {
549        (
550            Some(crate::trace::reduce_trace_axes(b, &right_trace)?),
551            false,
552        )
553    } else {
554        (None, crate::op_is_conj::<OpB>())
555    };
556    let b_view: StridedView<'_, T> = match b_buf.as_ref() {
557        Some(buf) => buf.view(),
558        None => StridedView::new(b.data(), b.dims(), b.strides(), b.offset())?,
559    };
560    let a_perm = a_view.permute(&plan.left_perm)?;
561    let b_perm = b_view.permute(&plan.right_perm)?;
562    let c_dims: Vec<usize> = plan
563        .c_to_internal_perm
564        .iter()
565        .map(|&axis| dest.dims()[axis])
566        .collect();
567    let c_strides: Vec<isize> = plan
568        .c_to_internal_perm
569        .iter()
570        .map(|&axis| dest.strides()[axis])
571        .collect();
572    let dest_offset = dest.offset();
573    let mut c_perm = RawStridedMut::new(dest.data_mut(), &c_dims, &c_strides, dest_offset)?;
574    let a_raw = RawStridedRef::new(
575        a_perm.data(),
576        a_perm.dims(),
577        a_perm.strides(),
578        a_perm.offset(),
579    )?;
580    let b_raw = RawStridedRef::new(
581        b_perm.data(),
582        b_perm.dims(),
583        b_perm.strides(),
584        b_perm.offset(),
585    )?;
586    let materialize = crate::make_conj_fn::<T>();
587    // BLAS has no conjugation flag. Materialize conjugation before preparing
588    // the raw backend operands, while retaining the same preflight contract.
589    if conj_a || conj_b {
590        let av = if conj_a {
591            let mut mapped =
592                unsafe { strided_view::StridedArray::<T>::col_major_uninit(a_perm.dims()) };
593            strided_kernel::map_into(&mut mapped.view_mut(), &a_perm, materialize.unwrap())?;
594            mapped
595        } else {
596            strided_view::StridedArray::from_parts(
597                a_perm.data().to_vec(),
598                a_perm.dims(),
599                a_perm.strides(),
600                a_perm.offset(),
601            )?
602        };
603        let bv = if conj_b {
604            let mut mapped =
605                unsafe { strided_view::StridedArray::<T>::col_major_uninit(b_perm.dims()) };
606            strided_kernel::map_into(&mut mapped.view_mut(), &b_perm, materialize.unwrap())?;
607            mapped
608        } else {
609            strided_view::StridedArray::from_parts(
610                b_perm.data().to_vec(),
611                b_perm.dims(),
612                b_perm.strides(),
613                b_perm.offset(),
614            )?
615        };
616        let ar = RawStridedRef::new(av.data(), av.dims(), av.strides(), av.view().offset())?;
617        let br = RawStridedRef::new(bv.data(), bv.dims(), bv.strides(), bv.view().offset())?;
618        return bgemm_raw_backend::<T, crate::backend::ActiveBackend>(
619            &mut c_perm,
620            &ar,
621            &br,
622            plan.batch.len(),
623            plan.lo.len(),
624            plan.ro.len(),
625            plan.sum.len(),
626            alpha,
627            ctx,
628        );
629    }
630    bgemm_raw_backend::<T, crate::backend::ActiveBackend>(
631        &mut c_perm,
632        &a_raw,
633        &b_raw,
634        plan.batch.len(),
635        plan.lo.len(),
636        plan.ro.len(),
637        plan.sum.len(),
638        alpha,
639        ctx,
640    )
641}
642
643/// Owned-input variant of [`einsum2_into_uninit`].
644#[allow(clippy::too_many_arguments)]
645#[cfg(not(any(feature = "blas", feature = "blas-inject")))]
646pub fn einsum2_into_owned_uninit<T, ID>(
647    dest: &mut RawStridedMut<'_, MaybeUninit<T>>,
648    a: strided_view::StridedArray<T>,
649    b: strided_view::StridedArray<T>,
650    ic: &[ID],
651    ia: &[ID],
652    ib: &[ID],
653    alpha: T,
654    ctx: &ExecContext,
655) -> Result<()>
656where
657    T: ScalarBase + strided_view::ElementOpApply,
658    ID: AxisId,
659{
660    einsum2_into_uninit(dest, &a.view(), &b.view(), ic, ia, ib, alpha, ctx)
661}
662
663/// Owned-input variant for BLAS-backed overwrite execution.
664#[allow(clippy::too_many_arguments)]
665#[cfg(any(feature = "blas", feature = "blas-inject"))]
666pub fn einsum2_into_owned_uninit<T, ID>(
667    dest: &mut RawStridedMut<'_, MaybeUninit<T>>,
668    a: strided_view::StridedArray<T>,
669    b: strided_view::StridedArray<T>,
670    ic: &[ID],
671    ia: &[ID],
672    ib: &[ID],
673    alpha: T,
674    ctx: &ExecContext,
675) -> Result<()>
676where
677    T: crate::Scalar,
678    ID: AxisId,
679{
680    einsum2_into_uninit(dest, &a.view(), &b.view(), ic, ia, ib, alpha, ctx)
681}
682
683/// Canonical raw overwrite-only GEMM entry point.
684#[allow(clippy::too_many_arguments)]
685pub fn bgemm_raw_strided_into_uninit<T>(
686    dest: &mut RawStridedMut<'_, MaybeUninit<T>>,
687    a: &RawStridedRef<'_, T>,
688    b: &RawStridedRef<'_, T>,
689    n_batch: usize,
690    n_lo: usize,
691    n_ro: usize,
692    n_sum: usize,
693    alpha: T,
694    ctx: &ExecContext,
695) -> Result<()>
696where
697    T: crate::Scalar,
698{
699    // This must precede label construction and all backend-specific
700    // materialization/allocation. It also validates shape agreement, output
701    // injectivity, conservative aliasing, checked products, and BLAS sizes.
702    let (groups, _, _, _) = preflight_raw_bgemm(dest, a, b, n_batch, n_lo, n_ro, n_sum)?;
703    let mut labels = Vec::with_capacity(groups.label_len);
704    labels.extend((0..groups.c_rank).map(|x| x));
705    #[cfg(not(any(feature = "blas", feature = "blas-inject")))]
706    let ic = labels[..groups.c_rank].to_vec();
707    #[cfg(not(any(feature = "blas", feature = "blas-inject")))]
708    let sum_start = groups.c_rank;
709    #[cfg(not(any(feature = "blas", feature = "blas-inject")))]
710    let ia = (0..n_lo)
711        .chain(sum_start..groups.label_len)
712        .chain(groups.c_ro_end..groups.c_rank)
713        .collect::<Vec<_>>();
714    #[cfg(not(any(feature = "blas", feature = "blas-inject")))]
715    let ib = (sum_start..groups.label_len)
716        .chain(n_lo..groups.c_ro_end)
717        .chain(groups.c_ro_end..groups.c_rank)
718        .collect::<Vec<_>>();
719    #[cfg(not(any(feature = "blas", feature = "blas-inject")))]
720    let av: StridedView<'_, T> =
721        unsafe { StridedView::new_unchecked(a.data(), a.dims(), a.strides(), a.offset()) };
722    #[cfg(not(any(feature = "blas", feature = "blas-inject")))]
723    let bv: StridedView<'_, T> =
724        unsafe { StridedView::new_unchecked(b.data(), b.dims(), b.strides(), b.offset()) };
725    #[cfg(any(feature = "blas", feature = "blas-inject"))]
726    {
727        return bgemm_raw_backend::<T, crate::backend::ActiveBackend>(
728            dest, a, b, n_batch, n_lo, n_ro, n_sum, alpha, ctx,
729        );
730    }
731    #[cfg(not(any(feature = "blas", feature = "blas-inject")))]
732    einsum2_into_uninit(dest, &av, &bv, &ic, &ia, &ib, alpha, ctx)
733}
734
735#[cfg(test)]
736mod tests {
737    use super::*;
738    use std::panic::{catch_unwind, AssertUnwindSafe};
739    use strided_view::StridedArray;
740
741    #[cfg(not(feature = "faer"))]
742    #[test]
743    fn matrix_product_writes_uninitialized_destination() {
744        let a = StridedArray::from_fn_row_major(&[2, 3], |idx| (idx[0] * 3 + idx[1] + 1) as f64);
745        let b = StridedArray::from_fn_row_major(&[3, 2], |idx| (idx[0] * 2 + idx[1] + 1) as f64);
746        let mut storage = vec![MaybeUninit::<f64>::uninit(); 4];
747        let dims = [2, 2];
748        let strides = [2, 1];
749        let mut c = RawStridedMut::new(&mut storage, &dims, &strides, 0).unwrap();
750        einsum2_into_uninit(
751            &mut c,
752            &a.view(),
753            &b.view(),
754            &['i', 'k'],
755            &['i', 'j'],
756            &['j', 'k'],
757            1.0,
758            &ExecContext::serial(),
759        )
760        .unwrap();
761        let values: Vec<f64> = storage
762            .into_iter()
763            .map(|x| unsafe { x.assume_init() })
764            .collect();
765        assert_eq!(values, vec![22.0, 28.0, 49.0, 64.0]);
766    }
767
768    #[test]
769    fn rejects_noninjective_destination_before_writing() {
770        let a = StridedArray::from_fn_col_major(&[2], |_| 1.0f64);
771        let b = StridedArray::from_fn_col_major(&[2], |_| 2.0f64);
772        let mut storage = vec![MaybeUninit::<f64>::uninit(); 1];
773        let dims = [2];
774        let strides = [0];
775        let mut c = RawStridedMut::new(&mut storage, &dims, &strides, 0).unwrap();
776        let err = einsum2_into_uninit(
777            &mut c,
778            &a.view(),
779            &b.view(),
780            &['i'],
781            &['i'],
782            &['i'],
783            1.0,
784            &ExecContext::serial(),
785        )
786        .unwrap_err();
787        assert!(matches!(err, crate::EinsumError::Strided(_)));
788    }
789
790    #[cfg(feature = "faer")]
791    #[test]
792    fn faer_uninit_gemm_reports_typed_unsupported_error() {
793        let a = StridedArray::from_fn_col_major(&[1, 1], |_| 1.0f64);
794        let b = StridedArray::from_fn_col_major(&[1, 1], |_| 1.0f64);
795        let mut storage = vec![MaybeUninit::<f64>::uninit()];
796        let dims = [1usize, 1];
797        let strides = [1isize, 1];
798        let mut c = RawStridedMut::new(&mut storage, &dims, &strides, 0).unwrap();
799        let err = einsum2_into_uninit(
800            &mut c,
801            &a.view(),
802            &b.view(),
803            &['i', 'k'],
804            &['i', 'j'],
805            &['j', 'k'],
806            1.0,
807            &ExecContext::serial(),
808        )
809        .unwrap_err();
810        assert!(matches!(err, crate::EinsumError::Unsupported(_)));
811    }
812
813    #[cfg(feature = "faer")]
814    #[test]
815    fn faer_uninit_gemm_validates_labels_before_backend_selection() {
816        let a = StridedArray::from_fn_col_major(&[2], |_| 1.0f64);
817        let b = StridedArray::from_fn_col_major(&[2], |_| 1.0f64);
818
819        let mut rank_storage = vec![MaybeUninit::<f64>::uninit(); 2];
820        let mut rank_dest = RawStridedMut::new(&mut rank_storage, &[2], &[1], 0).unwrap();
821        let rank_err = einsum2_into_uninit(
822            &mut rank_dest,
823            &a.view(),
824            &b.view(),
825            &['i'],
826            &['i', 'j'],
827            &['i'],
828            1.0,
829            &ExecContext::serial(),
830        )
831        .unwrap_err();
832        assert!(matches!(
833            rank_err,
834            crate::EinsumError::OutputShapeMismatch { .. }
835        ));
836
837        let mut shape_storage = vec![MaybeUninit::<f64>::uninit(); 3];
838        let mut shape_dest = RawStridedMut::new(&mut shape_storage, &[3], &[1], 0).unwrap();
839        let shape_err = einsum2_into_uninit(
840            &mut shape_dest,
841            &a.view(),
842            &b.view(),
843            &['i'],
844            &['i'],
845            &['i'],
846            1.0,
847            &ExecContext::serial(),
848        )
849        .unwrap_err();
850        assert!(matches!(
851            shape_err,
852            crate::EinsumError::OutputShapeMismatch { .. }
853        ));
854
855        let b_mismatched = StridedArray::from_fn_col_major(&[3], |_| 1.0f64);
856        let mut scalar_storage = vec![MaybeUninit::<f64>::uninit()];
857        let mut scalar_dest = RawStridedMut::new(&mut scalar_storage, &[], &[], 0).unwrap();
858        let dimension_err = einsum2_into_uninit(
859            &mut scalar_dest,
860            &a.view(),
861            &b_mismatched.view(),
862            &[],
863            &['i'],
864            &['i'],
865            1.0,
866            &ExecContext::serial(),
867        )
868        .unwrap_err();
869        assert!(matches!(
870            dimension_err,
871            crate::EinsumError::DimensionMismatch { .. }
872        ));
873    }
874
875    #[cfg(feature = "faer")]
876    #[test]
877    fn naive_overwrite_backend_covers_batches_and_conjugation() {
878        let a = StridedArray::from_fn_col_major(&[2, 2, 2], |idx| {
879            (1 + idx[0] + 2 * idx[1] + 4 * idx[2]) as f64
880        });
881        let b = StridedArray::from_fn_col_major(&[2, 2, 2], |idx| {
882            (1 + idx[0] + 2 * idx[1] + 4 * idx[2]) as f64
883        });
884        let a_raw = RawStridedRef::new(a.data(), a.dims(), a.strides(), a.view().offset()).unwrap();
885        let b_raw = RawStridedRef::new(b.data(), b.dims(), b.strides(), b.view().offset()).unwrap();
886        let a_op =
887            crate::contiguous::prepare_input_raw(&a_raw, 1, 1, true, false, false, None).unwrap();
888        let b_op =
889            crate::contiguous::prepare_input_raw(&b_raw, 1, 1, true, false, false, None).unwrap();
890
891        let mut storage = vec![MaybeUninit::<f64>::uninit(); 8];
892        let mut c = RawStridedMut::new(&mut storage, &[2, 2, 2], &[1, 2, 4], 0).unwrap();
893        let mut c_op = crate::contiguous::prepare_output_raw_uninit(&mut c, 1, 1, false).unwrap();
894        bgemm_contiguous_naive(
895            &mut c_op,
896            &a_op,
897            &b_op,
898            &[2],
899            2,
900            2,
901            2,
902            2.0,
903            &ExecContext::serial(),
904        )
905        .unwrap();
906        c_op.finalize().unwrap();
907
908        let values: Vec<f64> = storage
909            .into_iter()
910            .map(|x| unsafe { x.assume_init() })
911            .collect();
912        assert!(values.iter().all(|value| *value > 0.0));
913    }
914
915    #[cfg(not(feature = "faer"))]
916    #[test]
917    fn noncontiguous_output_is_written_back_after_overwrite() {
918        let a = StridedArray::from_fn_col_major(&[2, 2, 2], |idx| {
919            (1 + idx[0] + 2 * idx[1] + 4 * idx[2]) as f64
920        });
921        let b = StridedArray::from_fn_col_major(&[2, 2], |idx| (1 + idx[0] + 2 * idx[1]) as f64);
922        let mut storage = vec![MaybeUninit::<f64>::uninit(); 8];
923        let dims = [2usize, 2, 2];
924        let strides = [1isize, 4, 2];
925        let mut c = RawStridedMut::new(&mut storage, &dims, &strides, 0).unwrap();
926        einsum2_into_uninit(
927            &mut c,
928            &a.view(),
929            &b.view(),
930            &['i', 'j', 'k'],
931            &['i', 'j', 'l'],
932            &['l', 'k'],
933            1.0,
934            &ExecContext::serial(),
935        )
936        .unwrap();
937        let values: Vec<f64> = storage
938            .into_iter()
939            .map(|x| unsafe { x.assume_init() })
940            .collect();
941        assert_eq!(values, vec![11.0, 14.0, 23.0, 30.0, 17.0, 20.0, 37.0, 44.0]);
942    }
943
944    #[test]
945    fn raw_uninit_gemm_rejects_wrapping_group_partition_without_panicking() {
946        let a = StridedArray::from_fn_col_major(&[1], |_| 1.0f64);
947        let b = StridedArray::from_fn_col_major(&[1, 1], |_| 1.0f64);
948        let mut storage = vec![MaybeUninit::<f64>::uninit()];
949        let c_dims: [usize; 0] = [];
950        let c_strides: [isize; 0] = [];
951        let mut c = RawStridedMut::new(&mut storage, &c_dims, &c_strides, 0).unwrap();
952        let a_raw = RawStridedRef::new(a.data(), a.dims(), a.strides(), a.view().offset()).unwrap();
953        let b_raw = RawStridedRef::new(b.data(), b.dims(), b.strides(), b.view().offset()).unwrap();
954
955        let result = catch_unwind(AssertUnwindSafe(|| {
956            bgemm_raw_strided_into_uninit(
957                &mut c,
958                &a_raw,
959                &b_raw,
960                1,
961                usize::MAX,
962                0,
963                1,
964                1.0,
965                &ExecContext::serial(),
966            )
967        }));
968
969        assert!(result.is_ok(), "invalid group partition must not panic");
970        assert!(result.unwrap().is_err());
971    }
972}