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