1use crate::kernel::{
4 build_plan_fused, ensure_same_shape, for_each_inner_block_preordered, same_contiguous_layout,
5 sequential_contiguous_layout, total_len,
6};
7use crate::map_view::{map_into, zip_map2_into};
8use crate::maybe_sync::{MaybeSendSync, MaybeSync};
9use crate::reduce_view::reduce;
10use crate::simd;
11use crate::view::{StridedView, StridedViewMut};
12use crate::{Result, StridedError};
13use num_traits::{One, Zero};
14use std::ops::{Add, Mul};
15use strided_view::{ElementOp, ElementOpApply};
16
17#[cfg(feature = "parallel")]
18use crate::fuse::compute_costs;
19#[cfg(feature = "parallel")]
20use crate::threading::{
21 for_each_inner_block_with_offsets, mapreduce_threaded, SendPtr, MINTHREADLENGTH,
22};
23#[inline(always)]
33unsafe fn inner_loop_add<D: Copy + Add<S, Output = D>, S: Copy, Op: ElementOp<S>>(
34 dp: *mut D,
35 ds: isize,
36 sp: *const S,
37 ss: isize,
38 len: usize,
39) {
40 if ds == 1 && ss == 1 {
41 let dst = std::slice::from_raw_parts_mut(dp, len);
42 let src = std::slice::from_raw_parts(sp, len);
43 simd::dispatch_if_large(len, || {
44 for i in 0..len {
45 dst[i] = dst[i] + Op::apply(src[i]);
46 }
47 });
48 } else {
49 let mut dp = dp;
50 let mut sp = sp;
51 for _ in 0..len {
52 *dp = *dp + Op::apply(*sp);
53 dp = dp.offset(ds);
54 sp = sp.offset(ss);
55 }
56 }
57}
58
59#[inline(always)]
61unsafe fn inner_loop_mul<D: Copy + Mul<S, Output = D>, S: Copy, Op: ElementOp<S>>(
62 dp: *mut D,
63 ds: isize,
64 sp: *const S,
65 ss: isize,
66 len: usize,
67) {
68 if ds == 1 && ss == 1 {
69 let dst = std::slice::from_raw_parts_mut(dp, len);
70 let src = std::slice::from_raw_parts(sp, len);
71 simd::dispatch_if_large(len, || {
72 for i in 0..len {
73 dst[i] = dst[i] * Op::apply(src[i]);
74 }
75 });
76 } else {
77 let mut dp = dp;
78 let mut sp = sp;
79 for _ in 0..len {
80 *dp = *dp * Op::apply(*sp);
81 dp = dp.offset(ds);
82 sp = sp.offset(ss);
83 }
84 }
85}
86
87#[inline(always)]
89unsafe fn inner_loop_axpy<
90 D: Copy + Add<D, Output = D>,
91 S: Copy,
92 A: Copy + Mul<S, Output = D>,
93 Op: ElementOp<S>,
94>(
95 dp: *mut D,
96 ds: isize,
97 sp: *const S,
98 ss: isize,
99 len: usize,
100 alpha: A,
101) {
102 if ds == 1 && ss == 1 {
103 let dst = std::slice::from_raw_parts_mut(dp, len);
104 let src = std::slice::from_raw_parts(sp, len);
105 simd::dispatch_if_large(len, || {
106 for i in 0..len {
107 dst[i] = alpha * Op::apply(src[i]) + dst[i];
108 }
109 });
110 } else {
111 let mut dp = dp;
112 let mut sp = sp;
113 for _ in 0..len {
114 *dp = alpha * Op::apply(*sp) + *dp;
115 dp = dp.offset(ds);
116 sp = sp.offset(ss);
117 }
118 }
119}
120
121#[inline(always)]
123unsafe fn inner_loop_fma<
124 D: Copy + Add<D, Output = D>,
125 A: Copy + Mul<B, Output = D>,
126 B: Copy,
127 OpA: ElementOp<A>,
128 OpB: ElementOp<B>,
129>(
130 dp: *mut D,
131 ds: isize,
132 ap: *const A,
133 a_s: isize,
134 bp: *const B,
135 b_s: isize,
136 len: usize,
137) {
138 if ds == 1 && a_s == 1 && b_s == 1 {
139 let dst = std::slice::from_raw_parts_mut(dp, len);
140 let sa = std::slice::from_raw_parts(ap, len);
141 let sb = std::slice::from_raw_parts(bp, len);
142 simd::dispatch_if_large(len, || {
143 for i in 0..len {
144 dst[i] = dst[i] + OpA::apply(sa[i]) * OpB::apply(sb[i]);
145 }
146 });
147 } else {
148 let mut dp = dp;
149 let mut ap = ap;
150 let mut bp = bp;
151 for _ in 0..len {
152 *dp = *dp + OpA::apply(*ap) * OpB::apply(*bp);
153 dp = dp.offset(ds);
154 ap = ap.offset(a_s);
155 bp = bp.offset(b_s);
156 }
157 }
158}
159
160#[inline(always)]
162unsafe fn inner_loop_dot<
163 A: Copy + Mul<B, Output = R>,
164 B: Copy,
165 R: Copy + Add<R, Output = R>,
166 OpA: ElementOp<A>,
167 OpB: ElementOp<B>,
168>(
169 ap: *const A,
170 a_s: isize,
171 bp: *const B,
172 b_s: isize,
173 len: usize,
174 mut acc: R,
175) -> R {
176 if a_s == 1 && b_s == 1 {
177 let sa = std::slice::from_raw_parts(ap, len);
178 let sb = std::slice::from_raw_parts(bp, len);
179 simd::dispatch_if_large(len, || {
180 for i in 0..len {
181 acc = acc + OpA::apply(sa[i]) * OpB::apply(sb[i]);
182 }
183 });
184 } else {
185 let mut ap = ap;
186 let mut bp = bp;
187 for _ in 0..len {
188 acc = acc + OpA::apply(*ap) * OpB::apply(*bp);
189 ap = ap.offset(a_s);
190 bp = bp.offset(b_s);
191 }
192 }
193 acc
194}
195
196pub fn copy_into<T: Copy + MaybeSendSync, Op: ElementOp<T>>(
198 dest: &mut StridedViewMut<T>,
199 src: &StridedView<T, Op>,
200) -> Result<()> {
201 ensure_same_shape(dest.dims(), src.dims())?;
202
203 let dst_ptr = dest.as_mut_ptr();
204 let src_ptr = src.ptr();
205 let dst_dims = dest.dims();
206 let dst_strides = dest.strides();
207 let src_strides = src.strides();
208
209 if sequential_contiguous_layout(dst_dims, &[dst_strides, src_strides]).is_some() {
210 let len = total_len(dst_dims);
211 if Op::IS_IDENTITY {
212 debug_assert!(
213 {
214 let nbytes = len
215 .checked_mul(std::mem::size_of::<T>())
216 .expect("copy size must not overflow");
217 let dst_start = dst_ptr as usize;
218 let src_start = src_ptr as usize;
219 let dst_end = dst_start.saturating_add(nbytes);
220 let src_end = src_start.saturating_add(nbytes);
221 dst_end <= src_start || src_end <= dst_start
222 },
223 "overlapping src/dest is not supported"
224 );
225 unsafe { std::ptr::copy_nonoverlapping(src_ptr, dst_ptr, len) };
226 } else {
227 let dst = unsafe { std::slice::from_raw_parts_mut(dst_ptr, len) };
228 let src = unsafe { std::slice::from_raw_parts(src_ptr, len) };
229 simd::dispatch_if_large(len, || {
230 for i in 0..len {
231 dst[i] = Op::apply(src[i]);
232 }
233 });
234 }
235 return Ok(());
236 }
237
238 map_into(dest, src, |x| x)
239}
240
241pub fn copy_into_col_major<T: Copy + MaybeSendSync>(
245 dst: &mut StridedViewMut<T>,
246 src: &StridedView<T>,
247) -> Result<()> {
248 crate::threading::copy_into_col_major(dst, src)
249}
250
251pub fn add<
255 D: Copy + Add<S, Output = D> + MaybeSendSync,
256 S: Copy + MaybeSendSync,
257 Op: ElementOp<S>,
258>(
259 dest: &mut StridedViewMut<D>,
260 src: &StridedView<S, Op>,
261) -> Result<()> {
262 ensure_same_shape(dest.dims(), src.dims())?;
263
264 let dst_ptr = dest.as_mut_ptr();
265 let src_ptr = src.ptr();
266 let dst_dims = dest.dims();
267 let dst_strides = dest.strides();
268 let src_strides = src.strides();
269
270 if sequential_contiguous_layout(dst_dims, &[dst_strides, src_strides]).is_some() {
271 let len = total_len(dst_dims);
272 let dst = unsafe { std::slice::from_raw_parts_mut(dst_ptr, len) };
273 let src = unsafe { std::slice::from_raw_parts(src_ptr, len) };
274 simd::dispatch_if_large(len, || {
275 for i in 0..len {
276 dst[i] = dst[i] + Op::apply(src[i]);
277 }
278 });
279 return Ok(());
280 }
281
282 let strides_list: [&[isize]; 2] = [dst_strides, src_strides];
283 let elem_size = std::mem::size_of::<D>().max(std::mem::size_of::<S>());
284
285 let (fused_dims, ordered_strides, plan) =
286 build_plan_fused(dst_dims, &strides_list, Some(0), elem_size);
287
288 #[cfg(feature = "parallel")]
289 {
290 let total: usize = fused_dims.iter().product();
291 let nthreads = crate::execution_policy::rayon_threads();
292 if total > MINTHREADLENGTH && nthreads > 1 {
293 let dst_send = SendPtr(dst_ptr);
294 let src_send = SendPtr(src_ptr as *mut S);
295
296 let costs = compute_costs(&ordered_strides);
297 let initial_offsets = vec![0isize; strides_list.len()];
298 return mapreduce_threaded(
299 &fused_dims,
300 &plan.block,
301 &ordered_strides,
302 &initial_offsets,
303 &costs,
304 nthreads,
305 0,
306 1,
307 &|dims, blocks, strides_list, offsets| {
308 for_each_inner_block_with_offsets(
309 dims,
310 blocks,
311 strides_list,
312 offsets,
313 |offsets, len, strides| {
314 unsafe {
315 inner_loop_add::<D, S, Op>(
316 dst_send.as_ptr().offset(offsets[0]),
317 strides[0],
318 src_send.as_const().offset(offsets[1]),
319 strides[1],
320 len,
321 )
322 };
323 Ok(())
324 },
325 )
326 },
327 );
328 }
329 }
330
331 let initial_offsets = vec![0isize; ordered_strides.len()];
332 for_each_inner_block_preordered(
333 &fused_dims,
334 &plan.block,
335 &ordered_strides,
336 &initial_offsets,
337 |offsets, len, strides| {
338 unsafe {
339 inner_loop_add::<D, S, Op>(
340 dst_ptr.offset(offsets[0]),
341 strides[0],
342 src_ptr.offset(offsets[1]),
343 strides[1],
344 len,
345 )
346 };
347 Ok(())
348 },
349 )
350}
351
352pub fn mul<
356 D: Copy + Mul<S, Output = D> + MaybeSendSync,
357 S: Copy + MaybeSendSync,
358 Op: ElementOp<S>,
359>(
360 dest: &mut StridedViewMut<D>,
361 src: &StridedView<S, Op>,
362) -> Result<()> {
363 ensure_same_shape(dest.dims(), src.dims())?;
364
365 let dst_ptr = dest.as_mut_ptr();
366 let src_ptr = src.ptr();
367 let dst_dims = dest.dims();
368 let dst_strides = dest.strides();
369 let src_strides = src.strides();
370
371 if sequential_contiguous_layout(dst_dims, &[dst_strides, src_strides]).is_some() {
372 let len = total_len(dst_dims);
373 let dst = unsafe { std::slice::from_raw_parts_mut(dst_ptr, len) };
374 let src = unsafe { std::slice::from_raw_parts(src_ptr, len) };
375 simd::dispatch_if_large(len, || {
376 for i in 0..len {
377 dst[i] = dst[i] * Op::apply(src[i]);
378 }
379 });
380 return Ok(());
381 }
382
383 let strides_list: [&[isize]; 2] = [dst_strides, src_strides];
384 let elem_size = std::mem::size_of::<D>().max(std::mem::size_of::<S>());
385
386 let (fused_dims, ordered_strides, plan) =
387 build_plan_fused(dst_dims, &strides_list, Some(0), elem_size);
388
389 #[cfg(feature = "parallel")]
390 {
391 let total: usize = fused_dims.iter().product();
392 let nthreads = crate::execution_policy::rayon_threads();
393 if total > MINTHREADLENGTH && nthreads > 1 {
394 let dst_send = SendPtr(dst_ptr);
395 let src_send = SendPtr(src_ptr as *mut S);
396
397 let costs = compute_costs(&ordered_strides);
398 let initial_offsets = vec![0isize; strides_list.len()];
399 return mapreduce_threaded(
400 &fused_dims,
401 &plan.block,
402 &ordered_strides,
403 &initial_offsets,
404 &costs,
405 nthreads,
406 0,
407 1,
408 &|dims, blocks, strides_list, offsets| {
409 for_each_inner_block_with_offsets(
410 dims,
411 blocks,
412 strides_list,
413 offsets,
414 |offsets, len, strides| {
415 unsafe {
416 inner_loop_mul::<D, S, Op>(
417 dst_send.as_ptr().offset(offsets[0]),
418 strides[0],
419 src_send.as_const().offset(offsets[1]),
420 strides[1],
421 len,
422 )
423 };
424 Ok(())
425 },
426 )
427 },
428 );
429 }
430 }
431
432 let initial_offsets = vec![0isize; ordered_strides.len()];
433 for_each_inner_block_preordered(
434 &fused_dims,
435 &plan.block,
436 &ordered_strides,
437 &initial_offsets,
438 |offsets, len, strides| {
439 unsafe {
440 inner_loop_mul::<D, S, Op>(
441 dst_ptr.offset(offsets[0]),
442 strides[0],
443 src_ptr.offset(offsets[1]),
444 strides[1],
445 len,
446 )
447 };
448 Ok(())
449 },
450 )
451}
452
453pub fn axpy<D, S, A, Op>(
457 dest: &mut StridedViewMut<D>,
458 src: &StridedView<S, Op>,
459 alpha: A,
460) -> Result<()>
461where
462 A: Copy + Mul<S, Output = D> + MaybeSync,
463 D: Copy + Add<D, Output = D> + MaybeSendSync,
464 S: Copy + MaybeSendSync,
465 Op: ElementOp<S>,
466{
467 ensure_same_shape(dest.dims(), src.dims())?;
468
469 let dst_ptr = dest.as_mut_ptr();
470 let src_ptr = src.ptr();
471 let dst_dims = dest.dims();
472 let dst_strides = dest.strides();
473 let src_strides = src.strides();
474
475 if sequential_contiguous_layout(dst_dims, &[dst_strides, src_strides]).is_some() {
476 let len = total_len(dst_dims);
477 let dst = unsafe { std::slice::from_raw_parts_mut(dst_ptr, len) };
478 let src = unsafe { std::slice::from_raw_parts(src_ptr, len) };
479 simd::dispatch_if_large(len, || {
480 for i in 0..len {
481 dst[i] = alpha * Op::apply(src[i]) + dst[i];
482 }
483 });
484 return Ok(());
485 }
486
487 let strides_list: [&[isize]; 2] = [dst_strides, src_strides];
488 let elem_size = std::mem::size_of::<D>().max(std::mem::size_of::<S>());
489
490 let (fused_dims, ordered_strides, plan) =
491 build_plan_fused(dst_dims, &strides_list, Some(0), elem_size);
492
493 #[cfg(feature = "parallel")]
494 {
495 let total: usize = fused_dims.iter().product();
496 let nthreads = crate::execution_policy::rayon_threads();
497 if total > MINTHREADLENGTH && nthreads > 1 {
498 let dst_send = SendPtr(dst_ptr);
499 let src_send = SendPtr(src_ptr as *mut S);
500
501 let costs = compute_costs(&ordered_strides);
502 let initial_offsets = vec![0isize; strides_list.len()];
503 return mapreduce_threaded(
504 &fused_dims,
505 &plan.block,
506 &ordered_strides,
507 &initial_offsets,
508 &costs,
509 nthreads,
510 0,
511 1,
512 &|dims, blocks, strides_list, offsets| {
513 for_each_inner_block_with_offsets(
514 dims,
515 blocks,
516 strides_list,
517 offsets,
518 |offsets, len, strides| {
519 unsafe {
520 inner_loop_axpy::<D, S, A, Op>(
521 dst_send.as_ptr().offset(offsets[0]),
522 strides[0],
523 src_send.as_const().offset(offsets[1]),
524 strides[1],
525 len,
526 alpha,
527 )
528 };
529 Ok(())
530 },
531 )
532 },
533 );
534 }
535 }
536
537 let initial_offsets = vec![0isize; ordered_strides.len()];
538 for_each_inner_block_preordered(
539 &fused_dims,
540 &plan.block,
541 &ordered_strides,
542 &initial_offsets,
543 |offsets, len, strides| {
544 unsafe {
545 inner_loop_axpy::<D, S, A, Op>(
546 dst_ptr.offset(offsets[0]),
547 strides[0],
548 src_ptr.offset(offsets[1]),
549 strides[1],
550 len,
551 alpha,
552 )
553 };
554 Ok(())
555 },
556 )
557}
558
559pub fn fma<D, A, B, OpA, OpB>(
563 dest: &mut StridedViewMut<D>,
564 a: &StridedView<A, OpA>,
565 b: &StridedView<B, OpB>,
566) -> Result<()>
567where
568 A: Copy + Mul<B, Output = D> + MaybeSendSync,
569 B: Copy + MaybeSendSync,
570 D: Copy + Add<D, Output = D> + MaybeSendSync,
571 OpA: ElementOp<A>,
572 OpB: ElementOp<B>,
573{
574 ensure_same_shape(dest.dims(), a.dims())?;
575 ensure_same_shape(dest.dims(), b.dims())?;
576
577 let dst_ptr = dest.as_mut_ptr();
578 let a_ptr = a.ptr();
579 let b_ptr = b.ptr();
580 let dst_dims = dest.dims();
581 let dst_strides = dest.strides();
582 let a_strides = a.strides();
583 let b_strides = b.strides();
584
585 if sequential_contiguous_layout(dst_dims, &[dst_strides, a_strides, b_strides]).is_some() {
586 let len = total_len(dst_dims);
587 let dst = unsafe { std::slice::from_raw_parts_mut(dst_ptr, len) };
588 let sa = unsafe { std::slice::from_raw_parts(a_ptr, len) };
589 let sb = unsafe { std::slice::from_raw_parts(b_ptr, len) };
590 simd::dispatch_if_large(len, || {
591 for i in 0..len {
592 dst[i] = dst[i] + OpA::apply(sa[i]) * OpB::apply(sb[i]);
593 }
594 });
595 return Ok(());
596 }
597
598 let strides_list: [&[isize]; 3] = [dst_strides, a_strides, b_strides];
599 let elem_size = std::mem::size_of::<D>()
600 .max(std::mem::size_of::<A>())
601 .max(std::mem::size_of::<B>());
602
603 let (fused_dims, ordered_strides, plan) =
604 build_plan_fused(dst_dims, &strides_list, Some(0), elem_size);
605
606 #[cfg(feature = "parallel")]
607 {
608 let total: usize = fused_dims.iter().product();
609 let nthreads = crate::execution_policy::rayon_threads();
610 if total > MINTHREADLENGTH && nthreads > 1 {
611 let dst_send = SendPtr(dst_ptr);
612 let a_send = SendPtr(a_ptr as *mut A);
613 let b_send = SendPtr(b_ptr as *mut B);
614
615 let costs = compute_costs(&ordered_strides);
616 let initial_offsets = vec![0isize; strides_list.len()];
617 return mapreduce_threaded(
618 &fused_dims,
619 &plan.block,
620 &ordered_strides,
621 &initial_offsets,
622 &costs,
623 nthreads,
624 0,
625 1,
626 &|dims, blocks, strides_list, offsets| {
627 for_each_inner_block_with_offsets(
628 dims,
629 blocks,
630 strides_list,
631 offsets,
632 |offsets, len, strides| {
633 unsafe {
634 inner_loop_fma::<D, A, B, OpA, OpB>(
635 dst_send.as_ptr().offset(offsets[0]),
636 strides[0],
637 a_send.as_const().offset(offsets[1]),
638 strides[1],
639 b_send.as_const().offset(offsets[2]),
640 strides[2],
641 len,
642 )
643 };
644 Ok(())
645 },
646 )
647 },
648 );
649 }
650 }
651
652 let initial_offsets = vec![0isize; ordered_strides.len()];
653 for_each_inner_block_preordered(
654 &fused_dims,
655 &plan.block,
656 &ordered_strides,
657 &initial_offsets,
658 |offsets, len, strides| {
659 unsafe {
660 inner_loop_fma::<D, A, B, OpA, OpB>(
661 dst_ptr.offset(offsets[0]),
662 strides[0],
663 a_ptr.offset(offsets[1]),
664 strides[1],
665 b_ptr.offset(offsets[2]),
666 strides[2],
667 len,
668 )
669 };
670 Ok(())
671 },
672 )
673}
674
675#[cfg(feature = "parallel")]
676fn parallel_simd_sum<T: Copy + Zero + Add<Output = T> + simd::MaybeSimdOps + Send + Sync>(
677 src: &[T],
678) -> Option<T> {
679 if T::try_simd_sum(&[]).is_none() {
681 return None;
682 }
683 let nthreads = crate::execution_policy::rayon_threads();
684 let result = crate::threading::parallel_map_reduce(
685 0..src.len(),
686 nthreads,
687 &|range| T::try_simd_sum(&src[range]).unwrap(),
688 &|left, right| left + right,
689 );
690 Some(result)
691}
692
693pub fn sum<
695 T: Copy + Zero + Add<Output = T> + MaybeSendSync + simd::MaybeSimdOps,
696 Op: ElementOp<T>,
697>(
698 src: &StridedView<T, Op>,
699) -> Result<T> {
700 if Op::IS_IDENTITY {
702 if same_contiguous_layout(src.dims(), &[src.strides()]).is_some() {
703 let len = total_len(src.dims());
704 let src_slice = unsafe { std::slice::from_raw_parts(src.ptr(), len) };
705
706 #[cfg(feature = "parallel")]
707 if len > MINTHREADLENGTH {
708 if let Some(result) = parallel_simd_sum(src_slice) {
709 return Ok(result);
710 }
711 }
712
713 if let Some(result) = T::try_simd_sum(src_slice) {
714 return Ok(result);
715 }
716 }
717 }
718 reduce(src, |x| x, |a, b| a + b, T::zero())
719}
720
721pub fn dot<A, B, R, OpA, OpB>(a: &StridedView<A, OpA>, b: &StridedView<B, OpB>) -> Result<R>
726where
727 A: Copy + Mul<B, Output = R> + MaybeSendSync + 'static,
728 B: Copy + MaybeSendSync + 'static,
729 R: Copy + Zero + Add<Output = R> + MaybeSendSync + simd::MaybeSimdOps + 'static,
730 OpA: ElementOp<A>,
731 OpB: ElementOp<B>,
732{
733 ensure_same_shape(a.dims(), b.dims())?;
734
735 let a_ptr = a.ptr();
736 let b_ptr = b.ptr();
737 let a_strides = a.strides();
738 let b_strides = b.strides();
739 let a_dims = a.dims();
740
741 if same_contiguous_layout(a_dims, &[a_strides, b_strides]).is_some() {
742 let len = total_len(a_dims);
743
744 if OpA::IS_IDENTITY
746 && OpB::IS_IDENTITY
747 && std::any::TypeId::of::<A>() == std::any::TypeId::of::<R>()
748 && std::any::TypeId::of::<B>() == std::any::TypeId::of::<R>()
749 {
750 let sa = unsafe { std::slice::from_raw_parts(a_ptr as *const R, len) };
751 let sb = unsafe { std::slice::from_raw_parts(b_ptr as *const R, len) };
752 if let Some(result) = R::try_simd_dot(sa, sb) {
753 return Ok(result);
754 }
755 }
756
757 let sa = unsafe { std::slice::from_raw_parts(a_ptr, len) };
759 let sb = unsafe { std::slice::from_raw_parts(b_ptr, len) };
760 let mut acc = R::zero();
761 simd::dispatch_if_large(len, || {
762 for i in 0..len {
763 acc = acc + OpA::apply(sa[i]) * OpB::apply(sb[i]);
764 }
765 });
766 return Ok(acc);
767 }
768
769 let strides_list: [&[isize]; 2] = [a_strides, b_strides];
770 let elem_size = std::mem::size_of::<A>()
771 .max(std::mem::size_of::<B>())
772 .max(std::mem::size_of::<R>());
773
774 let (fused_dims, ordered_strides, plan) =
775 build_plan_fused(a_dims, &strides_list, None, elem_size);
776
777 let mut acc = R::zero();
778 let initial_offsets = vec![0isize; ordered_strides.len()];
779 for_each_inner_block_preordered(
780 &fused_dims,
781 &plan.block,
782 &ordered_strides,
783 &initial_offsets,
784 |offsets, len, strides| {
785 acc = unsafe {
786 inner_loop_dot::<A, B, R, OpA, OpB>(
787 a_ptr.offset(offsets[0]),
788 strides[0],
789 b_ptr.offset(offsets[1]),
790 strides[1],
791 len,
792 acc,
793 )
794 };
795 Ok(())
796 },
797 )?;
798
799 Ok(acc)
800}
801
802pub fn symmetrize_into<T>(dest: &mut StridedViewMut<T>, src: &StridedView<T>) -> Result<()>
804where
805 T: Copy
806 + Add<Output = T>
807 + Mul<Output = T>
808 + num_traits::FromPrimitive
809 + std::ops::Div<Output = T>
810 + MaybeSendSync,
811{
812 if src.ndim() != 2 {
813 return Err(StridedError::RankMismatch(src.ndim(), 2));
814 }
815 let rows = src.dims()[0];
816 let cols = src.dims()[1];
817 if rows != cols {
818 return Err(StridedError::NonSquare { rows, cols });
819 }
820
821 let src_t = src.permute(&[1, 0])?;
822 let half = T::from_f64(0.5).ok_or(StridedError::ScalarConversion)?;
823
824 zip_map2_into(dest, src, &src_t, |a, b| (a + b) * half)
825}
826
827pub fn symmetrize_conj_into<T>(dest: &mut StridedViewMut<T>, src: &StridedView<T>) -> Result<()>
829where
830 T: Copy
831 + ElementOpApply
832 + Add<Output = T>
833 + Mul<Output = T>
834 + num_traits::FromPrimitive
835 + std::ops::Div<Output = T>
836 + MaybeSendSync,
837{
838 if src.ndim() != 2 {
839 return Err(StridedError::RankMismatch(src.ndim(), 2));
840 }
841 let rows = src.dims()[0];
842 let cols = src.dims()[1];
843 if rows != cols {
844 return Err(StridedError::NonSquare { rows, cols });
845 }
846
847 let src_adj = src.adjoint_2d()?;
849 let half = T::from_f64(0.5).ok_or(StridedError::ScalarConversion)?;
850
851 zip_map2_into(dest, src, &src_adj, |a, b| (a + b) * half)
852}
853
854pub fn copy_scale<D, S, A, Op>(
858 dest: &mut StridedViewMut<D>,
859 src: &StridedView<S, Op>,
860 scale: A,
861) -> Result<()>
862where
863 A: Copy + Mul<S, Output = D> + MaybeSync,
864 D: Copy + MaybeSendSync,
865 S: Copy + MaybeSendSync,
866 Op: ElementOp<S>,
867{
868 map_into(dest, src, |x| scale * x)
869}
870
871pub fn copy_conj<T: Copy + ElementOpApply + MaybeSendSync>(
873 dest: &mut StridedViewMut<T>,
874 src: &StridedView<T>,
875) -> Result<()> {
876 let src_conj = src.conj();
877 copy_into(dest, &src_conj)
878}
879
880#[inline]
881fn element_transpose_is_identity<T: 'static>() -> bool {
882 use std::any::TypeId;
883
884 macro_rules! matches_type {
885 ($($ty:ty),* $(,)?) => {{
886 let id = TypeId::of::<T>();
887 false $(|| id == TypeId::of::<$ty>())*
888 }};
889 }
890
891 matches_type!(
892 f32,
893 f64,
894 i8,
895 i16,
896 i32,
897 i64,
898 i128,
899 isize,
900 u8,
901 u16,
902 u32,
903 u64,
904 u128,
905 usize,
906 num_complex::Complex32,
907 num_complex::Complex64,
908 )
909}
910
911#[inline]
912fn element_zero_is_all_bits_zero<T: 'static>() -> bool {
913 use std::any::TypeId;
914
915 macro_rules! matches_type {
916 ($($ty:ty),* $(,)?) => {{
917 let id = TypeId::of::<T>();
918 false $(|| id == TypeId::of::<$ty>())*
919 }};
920 }
921
922 matches_type!(
923 f32,
924 f64,
925 i8,
926 i16,
927 i32,
928 i64,
929 i128,
930 isize,
931 u8,
932 u16,
933 u32,
934 u64,
935 u128,
936 usize,
937 num_complex::Complex32,
938 num_complex::Complex64,
939 )
940}
941
942#[inline]
943unsafe fn fill_2d<T: Copy + MaybeSendSync>(
944 dst: *mut T,
945 dim0: usize,
946 dim1: usize,
947 dst_stride0: isize,
948 dst_stride1: isize,
949 value: T,
950) {
951 #[cfg(feature = "parallel")]
952 {
953 let total = dim0.saturating_mul(dim1);
954 let nthreads = crate::execution_policy::rayon_threads();
955 if total > MINTHREADLENGTH && nthreads > 1 {
956 let dst_send = SendPtr(dst);
957 if dst_stride0.unsigned_abs() <= dst_stride1.unsigned_abs() {
958 crate::threading::parallel_for_each(0..dim1, nthreads, &|columns| {
959 for j in columns {
960 let dst = dst_send.as_ptr();
961 unsafe {
962 let base = j as isize * dst_stride1;
963 for i in 0..dim0 {
964 *dst.offset(base + i as isize * dst_stride0) = value;
965 }
966 }
967 }
968 });
969 } else {
970 crate::threading::parallel_for_each(0..dim0, nthreads, &|rows| {
971 for i in rows {
972 let dst = dst_send.as_ptr();
973 unsafe {
974 let base = i as isize * dst_stride0;
975 for j in 0..dim1 {
976 *dst.offset(base + j as isize * dst_stride1) = value;
977 }
978 }
979 }
980 });
981 }
982 return;
983 }
984 }
985
986 if dst_stride0.unsigned_abs() <= dst_stride1.unsigned_abs() {
987 for j in 0..dim1 {
988 let base = j as isize * dst_stride1;
989 for i in 0..dim0 {
990 *dst.offset(base + i as isize * dst_stride0) = value;
991 }
992 }
993 } else {
994 for i in 0..dim0 {
995 let base = i as isize * dst_stride0;
996 for j in 0..dim1 {
997 *dst.offset(base + j as isize * dst_stride1) = value;
998 }
999 }
1000 }
1001}
1002
1003#[inline]
1004unsafe fn fill_contiguous<T>(dst: *mut T, len: usize, value: T)
1005where
1006 T: Copy + Zero + PartialEq + MaybeSendSync + 'static,
1007{
1008 if element_zero_is_all_bits_zero::<T>() && value == T::zero() {
1009 std::ptr::write_bytes(dst, 0, len);
1010 return;
1011 }
1012
1013 let dst = std::slice::from_raw_parts_mut(dst, len);
1014 dst.fill(value);
1015}
1016
1017#[inline(always)]
1018unsafe fn transpose_scale_4x4_f64(
1019 dst: *mut f64,
1020 dst_stride0: isize,
1021 dst_stride1: isize,
1022 src: *const f64,
1023 src_stride0: isize,
1024 src_stride1: isize,
1025 i: usize,
1026 j: usize,
1027 scale: f64,
1028) {
1029 let src_base = src.offset(i as isize * src_stride0 + j as isize * src_stride1);
1030
1031 let s00 = *src_base;
1032 let s10 = *src_base.offset(src_stride0);
1033 let s20 = *src_base.offset(2 * src_stride0);
1034 let s30 = *src_base.offset(3 * src_stride0);
1035
1036 let src_col1 = src_base.offset(src_stride1);
1037 let s01 = *src_col1;
1038 let s11 = *src_col1.offset(src_stride0);
1039 let s21 = *src_col1.offset(2 * src_stride0);
1040 let s31 = *src_col1.offset(3 * src_stride0);
1041
1042 let src_col2 = src_base.offset(2 * src_stride1);
1043 let s02 = *src_col2;
1044 let s12 = *src_col2.offset(src_stride0);
1045 let s22 = *src_col2.offset(2 * src_stride0);
1046 let s32 = *src_col2.offset(3 * src_stride0);
1047
1048 let src_col3 = src_base.offset(3 * src_stride1);
1049 let s03 = *src_col3;
1050 let s13 = *src_col3.offset(src_stride0);
1051 let s23 = *src_col3.offset(2 * src_stride0);
1052 let s33 = *src_col3.offset(3 * src_stride0);
1053
1054 let dst_row0 = dst.offset(j as isize * dst_stride0 + i as isize * dst_stride1);
1055 *dst_row0 = scale * s00;
1056 *dst_row0.offset(dst_stride0) = scale * s01;
1057 *dst_row0.offset(2 * dst_stride0) = scale * s02;
1058 *dst_row0.offset(3 * dst_stride0) = scale * s03;
1059
1060 let dst_row1 = dst_row0.offset(dst_stride1);
1061 *dst_row1 = scale * s10;
1062 *dst_row1.offset(dst_stride0) = scale * s11;
1063 *dst_row1.offset(2 * dst_stride0) = scale * s12;
1064 *dst_row1.offset(3 * dst_stride0) = scale * s13;
1065
1066 let dst_row2 = dst_row0.offset(2 * dst_stride1);
1067 *dst_row2 = scale * s20;
1068 *dst_row2.offset(dst_stride0) = scale * s21;
1069 *dst_row2.offset(2 * dst_stride0) = scale * s22;
1070 *dst_row2.offset(3 * dst_stride0) = scale * s23;
1071
1072 let dst_row3 = dst_row0.offset(3 * dst_stride1);
1073 *dst_row3 = scale * s30;
1074 *dst_row3.offset(dst_stride0) = scale * s31;
1075 *dst_row3.offset(2 * dst_stride0) = scale * s32;
1076 *dst_row3.offset(3 * dst_stride0) = scale * s33;
1077}
1078
1079#[inline]
1080unsafe fn copy_transpose_scale_2d_f64_tiled_raw(
1081 dst: *mut f64,
1082 dst_stride0: isize,
1083 dst_stride1: isize,
1084 src: *const f64,
1085 src_stride0: isize,
1086 src_stride1: isize,
1087 src_rows: usize,
1088 src_cols: usize,
1089 scale: f64,
1090) {
1091 const TILE: usize = 4;
1092 let row_full = src_rows / TILE * TILE;
1093 let col_full = src_cols / TILE * TILE;
1094
1095 #[cfg(feature = "parallel")]
1096 {
1097 let total = src_rows.saturating_mul(src_cols);
1098 let nthreads = crate::execution_policy::rayon_threads();
1099 if total > MINTHREADLENGTH && nthreads > 1 {
1100 let dst_send = SendPtr(dst);
1101 let src_send = SendPtr(src as *mut f64);
1102 let row_tiles = row_full / TILE;
1103 crate::threading::parallel_for_each(0..row_tiles, nthreads, &|tiles| {
1104 for tile_i in tiles {
1105 let i = tile_i * TILE;
1106 let dst = dst_send.as_ptr();
1107 let src = src_send.as_const();
1108 unsafe {
1109 let mut j = 0;
1110 while j < col_full {
1111 transpose_scale_4x4_f64(
1112 dst,
1113 dst_stride0,
1114 dst_stride1,
1115 src,
1116 src_stride0,
1117 src_stride1,
1118 i,
1119 j,
1120 scale,
1121 );
1122 j += TILE;
1123 }
1124 for j in col_full..src_cols {
1125 for ii in i..i + TILE {
1126 *dst.offset(j as isize * dst_stride0 + ii as isize * dst_stride1) =
1127 scale
1128 * *src.offset(
1129 ii as isize * src_stride0 + j as isize * src_stride1,
1130 );
1131 }
1132 }
1133 }
1134 }
1135 });
1136 for i in row_full..src_rows {
1137 for j in 0..src_cols {
1138 *dst.offset(j as isize * dst_stride0 + i as isize * dst_stride1) =
1139 scale * *src.offset(i as isize * src_stride0 + j as isize * src_stride1);
1140 }
1141 }
1142 return;
1143 }
1144 }
1145
1146 let mut i = 0;
1147 while i < row_full {
1148 let mut j = 0;
1149 while j < col_full {
1150 transpose_scale_4x4_f64(
1151 dst,
1152 dst_stride0,
1153 dst_stride1,
1154 src,
1155 src_stride0,
1156 src_stride1,
1157 i,
1158 j,
1159 scale,
1160 );
1161 j += TILE;
1162 }
1163 for j in col_full..src_cols {
1164 for ii in i..i + TILE {
1165 *dst.offset(j as isize * dst_stride0 + ii as isize * dst_stride1) =
1166 scale * *src.offset(ii as isize * src_stride0 + j as isize * src_stride1);
1167 }
1168 }
1169 i += TILE;
1170 }
1171 for i in row_full..src_rows {
1172 for j in 0..src_cols {
1173 *dst.offset(j as isize * dst_stride0 + i as isize * dst_stride1) =
1174 scale * *src.offset(i as isize * src_stride0 + j as isize * src_stride1);
1175 }
1176 }
1177}
1178
1179#[inline]
1180#[cfg(test)]
1181unsafe fn try_copy_transpose_scale_2d_f64_tiled(
1182 dest: &mut StridedViewMut<f64>,
1183 src: &StridedView<f64>,
1184 scale: f64,
1185) -> bool {
1186 if src.ndim() != 2 || dest.ndim() != 2 {
1187 return false;
1188 }
1189 let src_dims = src.dims();
1190 if dest.dims() != [src_dims[1], src_dims[0]] {
1191 return false;
1192 }
1193 if src.strides()[0] != 1 || dest.strides()[0] != 1 {
1194 return false;
1195 }
1196
1197 copy_transpose_scale_2d_f64_tiled_raw(
1198 dest.as_mut_ptr(),
1199 dest.strides()[0],
1200 dest.strides()[1],
1201 src.ptr(),
1202 src.strides()[0],
1203 src.strides()[1],
1204 src_dims[0],
1205 src_dims[1],
1206 scale,
1207 );
1208 true
1209}
1210
1211#[inline]
1212unsafe fn try_copy_transpose_scale_2d_f64_tiled_typed<T>(
1213 dest: &mut StridedViewMut<T>,
1214 src: &StridedView<T>,
1215 scale: T,
1216) -> bool
1217where
1218 T: Copy + 'static,
1219{
1220 if std::any::TypeId::of::<T>() != std::any::TypeId::of::<f64>() {
1221 return false;
1222 }
1223
1224 let scale = *(&scale as *const T).cast::<f64>();
1225 if src.ndim() != 2 || dest.ndim() != 2 {
1226 return false;
1227 }
1228 let src_dims = src.dims();
1229 if dest.dims() != [src_dims[1], src_dims[0]] {
1230 return false;
1231 }
1232 if src.strides()[0] != 1 || dest.strides()[0] != 1 {
1233 return false;
1234 }
1235
1236 copy_transpose_scale_2d_f64_tiled_raw(
1237 dest.as_mut_ptr().cast::<f64>(),
1238 dest.strides()[0],
1239 dest.strides()[1],
1240 src.ptr().cast::<f64>(),
1241 src.strides()[0],
1242 src.strides()[1],
1243 src_dims[0],
1244 src_dims[1],
1245 scale,
1246 );
1247 true
1248}
1249
1250#[inline(always)]
1251unsafe fn transpose_scale_4x4_identity<T>(
1252 dst: *mut T,
1253 dst_stride0: isize,
1254 dst_stride1: isize,
1255 src: *const T,
1256 src_stride0: isize,
1257 src_stride1: isize,
1258 i: usize,
1259 j: usize,
1260 scale: T,
1261) where
1262 T: Copy + Mul<Output = T>,
1263{
1264 let src_base = src.offset(i as isize * src_stride0 + j as isize * src_stride1);
1265
1266 let s00 = *src_base;
1267 let s10 = *src_base.offset(src_stride0);
1268 let s20 = *src_base.offset(2 * src_stride0);
1269 let s30 = *src_base.offset(3 * src_stride0);
1270
1271 let src_col1 = src_base.offset(src_stride1);
1272 let s01 = *src_col1;
1273 let s11 = *src_col1.offset(src_stride0);
1274 let s21 = *src_col1.offset(2 * src_stride0);
1275 let s31 = *src_col1.offset(3 * src_stride0);
1276
1277 let src_col2 = src_base.offset(2 * src_stride1);
1278 let s02 = *src_col2;
1279 let s12 = *src_col2.offset(src_stride0);
1280 let s22 = *src_col2.offset(2 * src_stride0);
1281 let s32 = *src_col2.offset(3 * src_stride0);
1282
1283 let src_col3 = src_base.offset(3 * src_stride1);
1284 let s03 = *src_col3;
1285 let s13 = *src_col3.offset(src_stride0);
1286 let s23 = *src_col3.offset(2 * src_stride0);
1287 let s33 = *src_col3.offset(3 * src_stride0);
1288
1289 let dst_row0 = dst.offset(j as isize * dst_stride0 + i as isize * dst_stride1);
1290 *dst_row0 = scale * s00;
1291 *dst_row0.offset(dst_stride0) = scale * s01;
1292 *dst_row0.offset(2 * dst_stride0) = scale * s02;
1293 *dst_row0.offset(3 * dst_stride0) = scale * s03;
1294
1295 let dst_row1 = dst_row0.offset(dst_stride1);
1296 *dst_row1 = scale * s10;
1297 *dst_row1.offset(dst_stride0) = scale * s11;
1298 *dst_row1.offset(2 * dst_stride0) = scale * s12;
1299 *dst_row1.offset(3 * dst_stride0) = scale * s13;
1300
1301 let dst_row2 = dst_row0.offset(2 * dst_stride1);
1302 *dst_row2 = scale * s20;
1303 *dst_row2.offset(dst_stride0) = scale * s21;
1304 *dst_row2.offset(2 * dst_stride0) = scale * s22;
1305 *dst_row2.offset(3 * dst_stride0) = scale * s23;
1306
1307 let dst_row3 = dst_row0.offset(3 * dst_stride1);
1308 *dst_row3 = scale * s30;
1309 *dst_row3.offset(dst_stride0) = scale * s31;
1310 *dst_row3.offset(2 * dst_stride0) = scale * s32;
1311 *dst_row3.offset(3 * dst_stride0) = scale * s33;
1312}
1313
1314#[inline]
1315unsafe fn copy_transpose_scale_2d_identity_tiled_raw<T>(
1316 dst: *mut T,
1317 dst_stride0: isize,
1318 dst_stride1: isize,
1319 src: *const T,
1320 src_stride0: isize,
1321 src_stride1: isize,
1322 src_rows: usize,
1323 src_cols: usize,
1324 scale: T,
1325) where
1326 T: Copy + Mul<Output = T> + MaybeSendSync,
1327{
1328 const TILE: usize = 4;
1329 let row_full = src_rows / TILE * TILE;
1330 let col_full = src_cols / TILE * TILE;
1331
1332 #[cfg(feature = "parallel")]
1333 {
1334 let total = src_rows.saturating_mul(src_cols);
1335 let nthreads = crate::execution_policy::rayon_threads();
1336 if total > MINTHREADLENGTH && nthreads > 1 {
1337 let dst_send = SendPtr(dst);
1338 let src_send = SendPtr(src as *mut T);
1339 let row_tiles = row_full / TILE;
1340 crate::threading::parallel_for_each(0..row_tiles, nthreads, &|tiles| {
1341 for tile_i in tiles {
1342 let i = tile_i * TILE;
1343 let dst = dst_send.as_ptr();
1344 let src = src_send.as_const();
1345 unsafe {
1346 let mut j = 0;
1347 while j < col_full {
1348 transpose_scale_4x4_identity(
1349 dst,
1350 dst_stride0,
1351 dst_stride1,
1352 src,
1353 src_stride0,
1354 src_stride1,
1355 i,
1356 j,
1357 scale,
1358 );
1359 j += TILE;
1360 }
1361 for j in col_full..src_cols {
1362 for ii in i..i + TILE {
1363 *dst.offset(j as isize * dst_stride0 + ii as isize * dst_stride1) =
1364 scale
1365 * *src.offset(
1366 ii as isize * src_stride0 + j as isize * src_stride1,
1367 );
1368 }
1369 }
1370 }
1371 }
1372 });
1373 for i in row_full..src_rows {
1374 for j in 0..src_cols {
1375 *dst.offset(j as isize * dst_stride0 + i as isize * dst_stride1) =
1376 scale * *src.offset(i as isize * src_stride0 + j as isize * src_stride1);
1377 }
1378 }
1379 return;
1380 }
1381 }
1382
1383 let mut i = 0;
1384 while i < row_full {
1385 let mut j = 0;
1386 while j < col_full {
1387 transpose_scale_4x4_identity(
1388 dst,
1389 dst_stride0,
1390 dst_stride1,
1391 src,
1392 src_stride0,
1393 src_stride1,
1394 i,
1395 j,
1396 scale,
1397 );
1398 j += TILE;
1399 }
1400 for j in col_full..src_cols {
1401 for ii in i..i + TILE {
1402 *dst.offset(j as isize * dst_stride0 + ii as isize * dst_stride1) =
1403 scale * *src.offset(ii as isize * src_stride0 + j as isize * src_stride1);
1404 }
1405 }
1406 i += TILE;
1407 }
1408 for i in row_full..src_rows {
1409 for j in 0..src_cols {
1410 *dst.offset(j as isize * dst_stride0 + i as isize * dst_stride1) =
1411 scale * *src.offset(i as isize * src_stride0 + j as isize * src_stride1);
1412 }
1413 }
1414}
1415
1416#[inline]
1417unsafe fn try_copy_transpose_scale_2d_identity_tiled<T>(
1418 dest: &mut StridedViewMut<T>,
1419 src: &StridedView<T>,
1420 scale: T,
1421) -> bool
1422where
1423 T: Copy + Mul<Output = T> + MaybeSendSync,
1424{
1425 if src.ndim() != 2 || dest.ndim() != 2 {
1426 return false;
1427 }
1428 let src_dims = src.dims();
1429 if dest.dims() != [src_dims[1], src_dims[0]] {
1430 return false;
1431 }
1432 if src.strides()[0] != 1 || dest.strides()[0] != 1 {
1433 return false;
1434 }
1435
1436 copy_transpose_scale_2d_identity_tiled_raw(
1437 dest.as_mut_ptr(),
1438 dest.strides()[0],
1439 dest.strides()[1],
1440 src.ptr(),
1441 src.strides()[0],
1442 src.strides()[1],
1443 src_dims[0],
1444 src_dims[1],
1445 scale,
1446 );
1447 true
1448}
1449
1450#[inline]
1451unsafe fn copy_transpose_scale_2d_loop<T>(
1452 dst: *mut T,
1453 dst_stride0: isize,
1454 dst_stride1: isize,
1455 src: *const T,
1456 src_stride0: isize,
1457 src_stride1: isize,
1458 src_rows: usize,
1459 src_cols: usize,
1460 scale: T,
1461) where
1462 T: Copy + ElementOpApply + Mul<Output = T> + MaybeSendSync,
1463{
1464 #[cfg(feature = "parallel")]
1465 {
1466 let total = src_rows.saturating_mul(src_cols);
1467 let nthreads = crate::execution_policy::rayon_threads();
1468 if total > MINTHREADLENGTH && nthreads > 1 {
1469 let dst_send = SendPtr(dst);
1470 let src_send = SendPtr(src as *mut T);
1471 if dst_stride0.unsigned_abs() <= dst_stride1.unsigned_abs() {
1472 crate::threading::parallel_for_each(0..src_rows, nthreads, &|rows| {
1473 for i in rows {
1474 let dst = dst_send.as_ptr();
1475 let src = src_send.as_const();
1476 unsafe {
1477 for j in 0..src_cols {
1478 let value = (*src
1479 .offset(i as isize * src_stride0 + j as isize * src_stride1))
1480 .transpose();
1481 *dst.offset(j as isize * dst_stride0 + i as isize * dst_stride1) =
1482 scale * value;
1483 }
1484 }
1485 }
1486 });
1487 } else {
1488 crate::threading::parallel_for_each(0..src_cols, nthreads, &|columns| {
1489 for j in columns {
1490 let dst = dst_send.as_ptr();
1491 let src = src_send.as_const();
1492 unsafe {
1493 for i in 0..src_rows {
1494 let value = (*src
1495 .offset(i as isize * src_stride0 + j as isize * src_stride1))
1496 .transpose();
1497 *dst.offset(j as isize * dst_stride0 + i as isize * dst_stride1) =
1498 scale * value;
1499 }
1500 }
1501 }
1502 });
1503 }
1504 return;
1505 }
1506 }
1507
1508 if dst_stride0.unsigned_abs() <= dst_stride1.unsigned_abs() {
1509 for i in 0..src_rows {
1510 for j in 0..src_cols {
1511 let value =
1512 (*src.offset(i as isize * src_stride0 + j as isize * src_stride1)).transpose();
1513 *dst.offset(j as isize * dst_stride0 + i as isize * dst_stride1) = scale * value;
1514 }
1515 }
1516 } else {
1517 for j in 0..src_cols {
1518 for i in 0..src_rows {
1519 let value =
1520 (*src.offset(i as isize * src_stride0 + j as isize * src_stride1)).transpose();
1521 *dst.offset(j as isize * dst_stride0 + i as isize * dst_stride1) = scale * value;
1522 }
1523 }
1524 }
1525}
1526
1527pub fn copy_transpose_scale_into<T>(
1529 dest: &mut StridedViewMut<T>,
1530 src: &StridedView<T>,
1531 scale: T,
1532) -> Result<()>
1533where
1534 T: Copy + ElementOpApply + Mul<Output = T> + Zero + One + PartialEq + MaybeSendSync + 'static,
1535{
1536 if src.ndim() != 2 || dest.ndim() != 2 {
1537 return Err(StridedError::RankMismatch(src.ndim(), 2));
1538 }
1539
1540 let src_dims = src.dims();
1541 let expected_dims = [src_dims[1], src_dims[0]];
1542 ensure_same_shape(dest.dims(), &expected_dims)?;
1543
1544 if scale == T::zero() {
1545 unsafe {
1546 if same_contiguous_layout(dest.dims(), &[dest.strides()]).is_some() {
1547 fill_contiguous(dest.as_mut_ptr(), total_len(dest.dims()), T::zero());
1548 } else {
1549 fill_2d(
1550 dest.as_mut_ptr(),
1551 dest.dims()[0],
1552 dest.dims()[1],
1553 dest.strides()[0],
1554 dest.strides()[1],
1555 T::zero(),
1556 );
1557 }
1558 }
1559 return Ok(());
1560 }
1561
1562 let transpose_is_identity = element_transpose_is_identity::<T>();
1563
1564 unsafe {
1565 if transpose_is_identity && try_copy_transpose_scale_2d_f64_tiled_typed(dest, src, scale) {
1566 return Ok(());
1567 }
1568 if transpose_is_identity && try_copy_transpose_scale_2d_identity_tiled(dest, src, scale) {
1569 return Ok(());
1570 }
1571 }
1572
1573 if scale == T::one() && transpose_is_identity {
1574 let src_t = src.permute(&[1, 0])?;
1575 #[cfg(feature = "parallel")]
1576 {
1577 return crate::threading::copy_permuted_with_active_policy(dest, &src_t);
1578 }
1579 #[cfg(not(feature = "parallel"))]
1580 return crate::threading::copy_permuted_serial(dest, &src_t);
1581 }
1582
1583 unsafe {
1584 copy_transpose_scale_2d_loop(
1585 dest.as_mut_ptr(),
1586 dest.strides()[0],
1587 dest.strides()[1],
1588 src.ptr(),
1589 src.strides()[0],
1590 src.strides()[1],
1591 src_dims[0],
1592 src_dims[1],
1593 scale,
1594 );
1595 }
1596 Ok(())
1597}
1598
1599#[cfg(test)]
1600mod tiled_tests {
1601 use super::*;
1602 use crate::view::StridedArray;
1603
1604 #[test]
1605 fn test_f64_tiled_transpose_scale_handles_remainders() {
1606 let rows = 7;
1607 let cols = 9;
1608 let a = StridedArray::<f64>::from_fn_col_major(&[rows, cols], |idx| {
1609 (idx[0] * 100 + idx[1]) as f64
1610 });
1611 let mut out = StridedArray::<f64>::col_major(&[cols, rows]);
1612
1613 let used_tiled = {
1614 let src = a.view();
1615 let mut dst = out.view_mut();
1616 unsafe { try_copy_transpose_scale_2d_f64_tiled(&mut dst, &src, 3.0) }
1617 };
1618
1619 assert!(used_tiled);
1620 for i in 0..rows {
1621 for j in 0..cols {
1622 assert_eq!(out.get(&[j, i]), 3.0 * a.get(&[i, j]));
1623 }
1624 }
1625 }
1626
1627 #[test]
1628 fn test_identity_tiled_transpose_scale_handles_integer_remainders() {
1629 let rows = 6;
1630 let cols = 5;
1631 let a = StridedArray::<u64>::from_fn_col_major(&[rows, cols], |idx| {
1632 (idx[0] * 100 + idx[1]) as u64
1633 });
1634 let mut out = StridedArray::<u64>::col_major(&[cols, rows]);
1635
1636 let used_tiled = {
1637 let src = a.view();
1638 let mut dst = out.view_mut();
1639 unsafe { try_copy_transpose_scale_2d_identity_tiled(&mut dst, &src, 2) }
1640 };
1641
1642 assert!(used_tiled);
1643 for i in 0..rows {
1644 for j in 0..cols {
1645 assert_eq!(out.get(&[j, i]), 2 * a.get(&[i, j]));
1646 }
1647 }
1648 }
1649
1650 #[cfg(feature = "parallel")]
1651 #[test]
1652 fn test_bounded_identity_tiled_transpose_uses_two_partitions_exactly_once() {
1653 use std::num::NonZeroUsize;
1654 use std::sync::atomic::{AtomicBool, AtomicUsize, Ordering};
1655 use std::time::{Duration, Instant};
1656
1657 struct TileState {
1658 active: AtomicUsize,
1659 max_active: AtomicUsize,
1660 released: AtomicBool,
1661 coverage: Box<[AtomicUsize]>,
1662 }
1663
1664 impl TileState {
1665 fn observe(&self, index: usize) {
1666 self.coverage[index].fetch_add(1, Ordering::SeqCst);
1667 let active = self.active.fetch_add(1, Ordering::SeqCst) + 1;
1668 self.max_active.fetch_max(active, Ordering::SeqCst);
1669 if active >= 2 {
1670 self.released.store(true, Ordering::Release);
1671 } else if !self.released.load(Ordering::Acquire) {
1672 let deadline = Instant::now() + Duration::from_secs(2);
1673 while !self.released.load(Ordering::Acquire) && Instant::now() < deadline {
1674 std::hint::spin_loop();
1675 }
1676 }
1677 self.active.fetch_sub(1, Ordering::SeqCst);
1678 }
1679 }
1680
1681 #[derive(Clone, Copy)]
1682 struct TrackedTile {
1683 index: usize,
1684 value: usize,
1685 state: &'static TileState,
1686 }
1687
1688 impl std::ops::Mul for TrackedTile {
1689 type Output = Self;
1690
1691 fn mul(self, rhs: Self) -> Self::Output {
1692 self.state.observe(rhs.index);
1693 Self {
1694 index: rhs.index,
1695 value: self.value * rhs.value,
1696 state: self.state,
1697 }
1698 }
1699 }
1700
1701 const ROWS: usize = 257;
1702 const COLS: usize = 129;
1703 const LEN: usize = ROWS * COLS;
1704 let state = Box::leak(Box::new(TileState {
1705 active: AtomicUsize::new(0),
1706 max_active: AtomicUsize::new(0),
1707 released: AtomicBool::new(false),
1708 coverage: (0..LEN)
1709 .map(|_| AtomicUsize::new(0))
1710 .collect::<Vec<_>>()
1711 .into_boxed_slice(),
1712 }));
1713 let source: Vec<_> = (0..LEN)
1714 .map(|index| TrackedTile {
1715 index,
1716 value: index + 1,
1717 state,
1718 })
1719 .collect();
1720 let mut destination = vec![
1721 TrackedTile {
1722 index: usize::MAX,
1723 value: 0,
1724 state,
1725 };
1726 LEN
1727 ];
1728 let scale = TrackedTile {
1729 index: usize::MAX,
1730 value: 3,
1731 state,
1732 };
1733 let two = NonZeroUsize::new(2).unwrap();
1734
1735 crate::threading::test_pool(4).install(|| {
1736 crate::with_execution_policy(
1737 crate::ExecutionPolicy::Rayon { max_threads: two },
1738 || unsafe {
1739 copy_transpose_scale_2d_identity_tiled_raw(
1740 destination.as_mut_ptr(),
1741 1,
1742 COLS as isize,
1743 source.as_ptr(),
1744 1,
1745 ROWS as isize,
1746 ROWS,
1747 COLS,
1748 scale,
1749 );
1750 },
1751 );
1752 });
1753
1754 assert_eq!(state.max_active.load(Ordering::SeqCst), 2);
1755 for (index, count) in state.coverage.iter().enumerate() {
1756 assert_eq!(
1757 count.load(Ordering::SeqCst),
1758 1,
1759 "tiled transpose source index {index} did not execute exactly once"
1760 );
1761 }
1762 for i in 0..ROWS {
1763 for j in 0..COLS {
1764 let source_index = i + ROWS * j;
1765 assert_eq!(destination[j + COLS * i].value, 3 * (source_index + 1));
1766 }
1767 }
1768 }
1769
1770 #[test]
1771 fn test_zero_scale_fills_non_contiguous_destination() {
1772 let rows = 3;
1773 let cols = 4;
1774 let a = StridedArray::<u64>::from_fn_col_major(&[rows, cols], |idx| {
1775 (idx[0] * 10 + idx[1] + 1) as u64
1776 });
1777 let mut out_base = StridedArray::<u64>::from_fn_col_major(&[rows, cols], |_| 99);
1778 let mut out_t = out_base.view_mut().permute(&[1, 0]).unwrap();
1779
1780 copy_transpose_scale_into(&mut out_t, &a.view(), 0).unwrap();
1781
1782 for i in 0..cols {
1783 for j in 0..rows {
1784 assert_eq!(out_t.get(&[i, j]), 0);
1785 }
1786 }
1787 }
1788}