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