1use 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#[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
223fn 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#[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#[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 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#[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#[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#[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 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}