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: usize = fused_dims.iter().product();
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: usize = fused_dims.iter().product();
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: usize = fused_dims.iter().product();
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: usize = fused_dims.iter().product();
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
806 if same_contiguous_layout(a_dims, &[a_strides, b_strides]).is_some() {
807 let len = total_len(a_dims);
808
809 if OpA::IS_IDENTITY
811 && OpB::IS_IDENTITY
812 && std::any::TypeId::of::<A>() == std::any::TypeId::of::<R>()
813 && std::any::TypeId::of::<B>() == std::any::TypeId::of::<R>()
814 {
815 let sa = unsafe { std::slice::from_raw_parts(a_ptr as *const R, len) };
816 let sb = unsafe { std::slice::from_raw_parts(b_ptr as *const R, len) };
817 if let Some(result) = R::try_simd_dot(sa, sb) {
818 return Ok(result);
819 }
820 }
821
822 let sa = unsafe { std::slice::from_raw_parts(a_ptr, len) };
824 let sb = unsafe { std::slice::from_raw_parts(b_ptr, len) };
825 let mut acc = R::zero();
826 simd::dispatch_if_large(len, || {
827 for i in 0..len {
828 acc = acc + OpA::apply(sa[i]) * OpB::apply(sb[i]);
829 }
830 });
831 return Ok(acc);
832 }
833
834 let strides_list: [&[isize]; 2] = [a_strides, b_strides];
835 let elem_size = std::mem::size_of::<A>()
836 .max(std::mem::size_of::<B>())
837 .max(std::mem::size_of::<R>());
838
839 let (fused_dims, ordered_strides, plan) =
840 build_plan_fused(a_dims, &strides_list, None, elem_size);
841
842 let mut acc = R::zero();
843 let initial_offsets = vec![0isize; ordered_strides.len()];
844 for_each_inner_block_preordered(
845 &fused_dims,
846 &plan.block,
847 &ordered_strides,
848 &initial_offsets,
849 |offsets, len, strides| {
850 acc = unsafe {
851 inner_loop_dot::<A, B, R, OpA, OpB>(
852 a_ptr.offset(offsets[0]),
853 strides[0],
854 b_ptr.offset(offsets[1]),
855 strides[1],
856 len,
857 acc,
858 )
859 };
860 Ok(())
861 },
862 )?;
863
864 Ok(acc)
865}
866
867pub fn symmetrize_into<T>(dest: &mut StridedViewMut<T>, src: &StridedView<T>) -> Result<()>
869where
870 T: Copy
871 + Add<Output = T>
872 + Mul<Output = T>
873 + num_traits::FromPrimitive
874 + std::ops::Div<Output = T>
875 + MaybeSendSync,
876{
877 if src.ndim() != 2 {
878 return Err(StridedError::RankMismatch(src.ndim(), 2));
879 }
880 let rows = src.dims()[0];
881 let cols = src.dims()[1];
882 if rows != cols {
883 return Err(StridedError::NonSquare { rows, cols });
884 }
885
886 let src_t = src.permute(&[1, 0])?;
887 let half = T::from_f64(0.5).ok_or(StridedError::ScalarConversion)?;
888
889 zip_map2_into(dest, src, &src_t, |a, b| (a + b) * half)
890}
891
892pub fn symmetrize_conj_into<T>(dest: &mut StridedViewMut<T>, src: &StridedView<T>) -> Result<()>
894where
895 T: Copy
896 + ElementOpApply
897 + Add<Output = T>
898 + Mul<Output = T>
899 + num_traits::FromPrimitive
900 + std::ops::Div<Output = T>
901 + MaybeSendSync,
902{
903 if src.ndim() != 2 {
904 return Err(StridedError::RankMismatch(src.ndim(), 2));
905 }
906 let rows = src.dims()[0];
907 let cols = src.dims()[1];
908 if rows != cols {
909 return Err(StridedError::NonSquare { rows, cols });
910 }
911
912 let src_adj = src.adjoint_2d()?;
914 let half = T::from_f64(0.5).ok_or(StridedError::ScalarConversion)?;
915
916 zip_map2_into(dest, src, &src_adj, |a, b| (a + b) * half)
917}
918
919pub fn copy_scale<D, S, A, Op>(
923 dest: &mut StridedViewMut<D>,
924 src: &StridedView<S, Op>,
925 scale: A,
926) -> Result<()>
927where
928 A: Copy + Mul<S, Output = D> + MaybeSync,
929 D: Copy + MaybeSendSync,
930 S: Copy + MaybeSendSync,
931 Op: ElementOp<S>,
932{
933 map_into(dest, src, |x| scale * x)
934}
935
936pub fn copy_conj<T: Copy + ElementOpApply + MaybeSendSync>(
938 dest: &mut StridedViewMut<T>,
939 src: &StridedView<T>,
940) -> Result<()> {
941 let src_conj = src.conj();
942 copy_into(dest, &src_conj)
943}
944
945#[inline]
946fn element_transpose_is_identity<T: 'static>() -> bool {
947 use std::any::TypeId;
948
949 macro_rules! matches_type {
950 ($($ty:ty),* $(,)?) => {{
951 let id = TypeId::of::<T>();
952 false $(|| id == TypeId::of::<$ty>())*
953 }};
954 }
955
956 matches_type!(
957 f32,
958 f64,
959 i8,
960 i16,
961 i32,
962 i64,
963 i128,
964 isize,
965 u8,
966 u16,
967 u32,
968 u64,
969 u128,
970 usize,
971 num_complex::Complex32,
972 num_complex::Complex64,
973 )
974}
975
976#[inline]
977fn element_zero_is_all_bits_zero<T: 'static>() -> bool {
978 use std::any::TypeId;
979
980 macro_rules! matches_type {
981 ($($ty:ty),* $(,)?) => {{
982 let id = TypeId::of::<T>();
983 false $(|| id == TypeId::of::<$ty>())*
984 }};
985 }
986
987 matches_type!(
988 f32,
989 f64,
990 i8,
991 i16,
992 i32,
993 i64,
994 i128,
995 isize,
996 u8,
997 u16,
998 u32,
999 u64,
1000 u128,
1001 usize,
1002 num_complex::Complex32,
1003 num_complex::Complex64,
1004 )
1005}
1006
1007#[inline]
1008unsafe fn fill_2d<T: Copy + MaybeSendSync>(
1009 dst: *mut T,
1010 dim0: usize,
1011 dim1: usize,
1012 dst_stride0: isize,
1013 dst_stride1: isize,
1014 value: T,
1015) {
1016 #[cfg(feature = "parallel")]
1017 {
1018 let total = dim0.saturating_mul(dim1);
1019 let nthreads = crate::execution_policy::rayon_threads();
1020 if total > MINTHREADLENGTH && nthreads > 1 {
1021 let dst_send = SendPtr(dst);
1022 if dst_stride0.unsigned_abs() <= dst_stride1.unsigned_abs() {
1023 crate::threading::parallel_for_each(0..dim1, nthreads, &|columns| {
1024 for j in columns {
1025 let dst = dst_send.as_ptr();
1026 unsafe {
1027 let base = j as isize * dst_stride1;
1028 for i in 0..dim0 {
1029 *dst.offset(base + i as isize * dst_stride0) = value;
1030 }
1031 }
1032 }
1033 });
1034 } else {
1035 crate::threading::parallel_for_each(0..dim0, nthreads, &|rows| {
1036 for i in rows {
1037 let dst = dst_send.as_ptr();
1038 unsafe {
1039 let base = i as isize * dst_stride0;
1040 for j in 0..dim1 {
1041 *dst.offset(base + j as isize * dst_stride1) = value;
1042 }
1043 }
1044 }
1045 });
1046 }
1047 return;
1048 }
1049 }
1050
1051 if dst_stride0.unsigned_abs() <= dst_stride1.unsigned_abs() {
1052 for j in 0..dim1 {
1053 let base = j as isize * dst_stride1;
1054 for i in 0..dim0 {
1055 *dst.offset(base + i as isize * dst_stride0) = value;
1056 }
1057 }
1058 } else {
1059 for i in 0..dim0 {
1060 let base = i as isize * dst_stride0;
1061 for j in 0..dim1 {
1062 *dst.offset(base + j as isize * dst_stride1) = value;
1063 }
1064 }
1065 }
1066}
1067
1068#[inline]
1069unsafe fn fill_contiguous<T>(dst: *mut T, len: usize, value: T)
1070where
1071 T: Copy + Zero + PartialEq + MaybeSendSync + 'static,
1072{
1073 if element_zero_is_all_bits_zero::<T>() && value == T::zero() {
1074 std::ptr::write_bytes(dst, 0, len);
1075 return;
1076 }
1077
1078 let dst = std::slice::from_raw_parts_mut(dst, len);
1079 dst.fill(value);
1080}
1081
1082#[inline(always)]
1083unsafe fn transpose_scale_4x4_f64(
1084 dst: *mut f64,
1085 dst_stride0: isize,
1086 dst_stride1: isize,
1087 src: *const f64,
1088 src_stride0: isize,
1089 src_stride1: isize,
1090 i: usize,
1091 j: usize,
1092 scale: f64,
1093) {
1094 let src_base = src.offset(i as isize * src_stride0 + j as isize * src_stride1);
1095
1096 let s00 = *src_base;
1097 let s10 = *src_base.offset(src_stride0);
1098 let s20 = *src_base.offset(2 * src_stride0);
1099 let s30 = *src_base.offset(3 * src_stride0);
1100
1101 let src_col1 = src_base.offset(src_stride1);
1102 let s01 = *src_col1;
1103 let s11 = *src_col1.offset(src_stride0);
1104 let s21 = *src_col1.offset(2 * src_stride0);
1105 let s31 = *src_col1.offset(3 * src_stride0);
1106
1107 let src_col2 = src_base.offset(2 * src_stride1);
1108 let s02 = *src_col2;
1109 let s12 = *src_col2.offset(src_stride0);
1110 let s22 = *src_col2.offset(2 * src_stride0);
1111 let s32 = *src_col2.offset(3 * src_stride0);
1112
1113 let src_col3 = src_base.offset(3 * src_stride1);
1114 let s03 = *src_col3;
1115 let s13 = *src_col3.offset(src_stride0);
1116 let s23 = *src_col3.offset(2 * src_stride0);
1117 let s33 = *src_col3.offset(3 * src_stride0);
1118
1119 let dst_row0 = dst.offset(j as isize * dst_stride0 + i as isize * dst_stride1);
1120 *dst_row0 = scale * s00;
1121 *dst_row0.offset(dst_stride0) = scale * s01;
1122 *dst_row0.offset(2 * dst_stride0) = scale * s02;
1123 *dst_row0.offset(3 * dst_stride0) = scale * s03;
1124
1125 let dst_row1 = dst_row0.offset(dst_stride1);
1126 *dst_row1 = scale * s10;
1127 *dst_row1.offset(dst_stride0) = scale * s11;
1128 *dst_row1.offset(2 * dst_stride0) = scale * s12;
1129 *dst_row1.offset(3 * dst_stride0) = scale * s13;
1130
1131 let dst_row2 = dst_row0.offset(2 * dst_stride1);
1132 *dst_row2 = scale * s20;
1133 *dst_row2.offset(dst_stride0) = scale * s21;
1134 *dst_row2.offset(2 * dst_stride0) = scale * s22;
1135 *dst_row2.offset(3 * dst_stride0) = scale * s23;
1136
1137 let dst_row3 = dst_row0.offset(3 * dst_stride1);
1138 *dst_row3 = scale * s30;
1139 *dst_row3.offset(dst_stride0) = scale * s31;
1140 *dst_row3.offset(2 * dst_stride0) = scale * s32;
1141 *dst_row3.offset(3 * dst_stride0) = scale * s33;
1142}
1143
1144#[inline]
1145unsafe fn copy_transpose_scale_2d_f64_tiled_raw(
1146 dst: *mut f64,
1147 dst_stride0: isize,
1148 dst_stride1: isize,
1149 src: *const f64,
1150 src_stride0: isize,
1151 src_stride1: isize,
1152 src_rows: usize,
1153 src_cols: usize,
1154 scale: f64,
1155) {
1156 const TILE: usize = 4;
1157 let row_full = src_rows / TILE * TILE;
1158 let col_full = src_cols / TILE * TILE;
1159
1160 #[cfg(feature = "parallel")]
1161 {
1162 let total = src_rows.saturating_mul(src_cols);
1163 let nthreads = crate::execution_policy::rayon_threads();
1164 if total > MINTHREADLENGTH && nthreads > 1 {
1165 let dst_send = SendPtr(dst);
1166 let src_send = SendPtr(src as *mut f64);
1167 let row_tiles = row_full / TILE;
1168 crate::threading::parallel_for_each(0..row_tiles, nthreads, &|tiles| {
1169 for tile_i in tiles {
1170 let i = tile_i * TILE;
1171 let dst = dst_send.as_ptr();
1172 let src = src_send.as_const();
1173 unsafe {
1174 let mut j = 0;
1175 while j < col_full {
1176 transpose_scale_4x4_f64(
1177 dst,
1178 dst_stride0,
1179 dst_stride1,
1180 src,
1181 src_stride0,
1182 src_stride1,
1183 i,
1184 j,
1185 scale,
1186 );
1187 j += TILE;
1188 }
1189 for j in col_full..src_cols {
1190 for ii in i..i + TILE {
1191 *dst.offset(j as isize * dst_stride0 + ii as isize * dst_stride1) =
1192 scale
1193 * *src.offset(
1194 ii as isize * src_stride0 + j as isize * src_stride1,
1195 );
1196 }
1197 }
1198 }
1199 }
1200 });
1201 for i in row_full..src_rows {
1202 for j in 0..src_cols {
1203 *dst.offset(j as isize * dst_stride0 + i as isize * dst_stride1) =
1204 scale * *src.offset(i as isize * src_stride0 + j as isize * src_stride1);
1205 }
1206 }
1207 return;
1208 }
1209 }
1210
1211 let mut i = 0;
1212 while i < row_full {
1213 let mut j = 0;
1214 while j < col_full {
1215 transpose_scale_4x4_f64(
1216 dst,
1217 dst_stride0,
1218 dst_stride1,
1219 src,
1220 src_stride0,
1221 src_stride1,
1222 i,
1223 j,
1224 scale,
1225 );
1226 j += TILE;
1227 }
1228 for j in col_full..src_cols {
1229 for ii in i..i + TILE {
1230 *dst.offset(j as isize * dst_stride0 + ii as isize * dst_stride1) =
1231 scale * *src.offset(ii as isize * src_stride0 + j as isize * src_stride1);
1232 }
1233 }
1234 i += TILE;
1235 }
1236 for i in row_full..src_rows {
1237 for j in 0..src_cols {
1238 *dst.offset(j as isize * dst_stride0 + i as isize * dst_stride1) =
1239 scale * *src.offset(i as isize * src_stride0 + j as isize * src_stride1);
1240 }
1241 }
1242}
1243
1244#[inline]
1245#[cfg(test)]
1246unsafe fn try_copy_transpose_scale_2d_f64_tiled(
1247 dest: &mut StridedViewMut<f64>,
1248 src: &StridedView<f64>,
1249 scale: f64,
1250) -> bool {
1251 if src.ndim() != 2 || dest.ndim() != 2 {
1252 return false;
1253 }
1254 let src_dims = src.dims();
1255 if dest.dims() != [src_dims[1], src_dims[0]] {
1256 return false;
1257 }
1258 if src.strides()[0] != 1 || dest.strides()[0] != 1 {
1259 return false;
1260 }
1261
1262 copy_transpose_scale_2d_f64_tiled_raw(
1263 dest.as_mut_ptr(),
1264 dest.strides()[0],
1265 dest.strides()[1],
1266 src.ptr(),
1267 src.strides()[0],
1268 src.strides()[1],
1269 src_dims[0],
1270 src_dims[1],
1271 scale,
1272 );
1273 true
1274}
1275
1276#[inline]
1277unsafe fn try_copy_transpose_scale_2d_f64_tiled_typed<T>(
1278 dest: &mut StridedViewMut<T>,
1279 src: &StridedView<T>,
1280 scale: T,
1281) -> bool
1282where
1283 T: Copy + 'static,
1284{
1285 if std::any::TypeId::of::<T>() != std::any::TypeId::of::<f64>() {
1286 return false;
1287 }
1288
1289 let scale = *(&scale as *const T).cast::<f64>();
1290 if src.ndim() != 2 || dest.ndim() != 2 {
1291 return false;
1292 }
1293 let src_dims = src.dims();
1294 if dest.dims() != [src_dims[1], src_dims[0]] {
1295 return false;
1296 }
1297 if src.strides()[0] != 1 || dest.strides()[0] != 1 {
1298 return false;
1299 }
1300
1301 copy_transpose_scale_2d_f64_tiled_raw(
1302 dest.as_mut_ptr().cast::<f64>(),
1303 dest.strides()[0],
1304 dest.strides()[1],
1305 src.ptr().cast::<f64>(),
1306 src.strides()[0],
1307 src.strides()[1],
1308 src_dims[0],
1309 src_dims[1],
1310 scale,
1311 );
1312 true
1313}
1314
1315#[inline(always)]
1316unsafe fn transpose_scale_4x4_identity<T>(
1317 dst: *mut T,
1318 dst_stride0: isize,
1319 dst_stride1: isize,
1320 src: *const T,
1321 src_stride0: isize,
1322 src_stride1: isize,
1323 i: usize,
1324 j: usize,
1325 scale: T,
1326) where
1327 T: Copy + Mul<Output = T>,
1328{
1329 let src_base = src.offset(i as isize * src_stride0 + j as isize * src_stride1);
1330
1331 let s00 = *src_base;
1332 let s10 = *src_base.offset(src_stride0);
1333 let s20 = *src_base.offset(2 * src_stride0);
1334 let s30 = *src_base.offset(3 * src_stride0);
1335
1336 let src_col1 = src_base.offset(src_stride1);
1337 let s01 = *src_col1;
1338 let s11 = *src_col1.offset(src_stride0);
1339 let s21 = *src_col1.offset(2 * src_stride0);
1340 let s31 = *src_col1.offset(3 * src_stride0);
1341
1342 let src_col2 = src_base.offset(2 * src_stride1);
1343 let s02 = *src_col2;
1344 let s12 = *src_col2.offset(src_stride0);
1345 let s22 = *src_col2.offset(2 * src_stride0);
1346 let s32 = *src_col2.offset(3 * src_stride0);
1347
1348 let src_col3 = src_base.offset(3 * src_stride1);
1349 let s03 = *src_col3;
1350 let s13 = *src_col3.offset(src_stride0);
1351 let s23 = *src_col3.offset(2 * src_stride0);
1352 let s33 = *src_col3.offset(3 * src_stride0);
1353
1354 let dst_row0 = dst.offset(j as isize * dst_stride0 + i as isize * dst_stride1);
1355 *dst_row0 = scale * s00;
1356 *dst_row0.offset(dst_stride0) = scale * s01;
1357 *dst_row0.offset(2 * dst_stride0) = scale * s02;
1358 *dst_row0.offset(3 * dst_stride0) = scale * s03;
1359
1360 let dst_row1 = dst_row0.offset(dst_stride1);
1361 *dst_row1 = scale * s10;
1362 *dst_row1.offset(dst_stride0) = scale * s11;
1363 *dst_row1.offset(2 * dst_stride0) = scale * s12;
1364 *dst_row1.offset(3 * dst_stride0) = scale * s13;
1365
1366 let dst_row2 = dst_row0.offset(2 * dst_stride1);
1367 *dst_row2 = scale * s20;
1368 *dst_row2.offset(dst_stride0) = scale * s21;
1369 *dst_row2.offset(2 * dst_stride0) = scale * s22;
1370 *dst_row2.offset(3 * dst_stride0) = scale * s23;
1371
1372 let dst_row3 = dst_row0.offset(3 * dst_stride1);
1373 *dst_row3 = scale * s30;
1374 *dst_row3.offset(dst_stride0) = scale * s31;
1375 *dst_row3.offset(2 * dst_stride0) = scale * s32;
1376 *dst_row3.offset(3 * dst_stride0) = scale * s33;
1377}
1378
1379#[inline]
1380unsafe fn copy_transpose_scale_2d_identity_tiled_raw<T>(
1381 dst: *mut T,
1382 dst_stride0: isize,
1383 dst_stride1: isize,
1384 src: *const T,
1385 src_stride0: isize,
1386 src_stride1: isize,
1387 src_rows: usize,
1388 src_cols: usize,
1389 scale: T,
1390) where
1391 T: Copy + Mul<Output = T> + MaybeSendSync,
1392{
1393 const TILE: usize = 4;
1394 let row_full = src_rows / TILE * TILE;
1395 let col_full = src_cols / TILE * TILE;
1396
1397 #[cfg(feature = "parallel")]
1398 {
1399 let total = src_rows.saturating_mul(src_cols);
1400 let nthreads = crate::execution_policy::rayon_threads();
1401 if total > MINTHREADLENGTH && nthreads > 1 {
1402 let dst_send = SendPtr(dst);
1403 let src_send = SendPtr(src as *mut T);
1404 let row_tiles = row_full / TILE;
1405 crate::threading::parallel_for_each(0..row_tiles, nthreads, &|tiles| {
1406 for tile_i in tiles {
1407 let i = tile_i * TILE;
1408 let dst = dst_send.as_ptr();
1409 let src = src_send.as_const();
1410 unsafe {
1411 let mut j = 0;
1412 while j < col_full {
1413 transpose_scale_4x4_identity(
1414 dst,
1415 dst_stride0,
1416 dst_stride1,
1417 src,
1418 src_stride0,
1419 src_stride1,
1420 i,
1421 j,
1422 scale,
1423 );
1424 j += TILE;
1425 }
1426 for j in col_full..src_cols {
1427 for ii in i..i + TILE {
1428 *dst.offset(j as isize * dst_stride0 + ii as isize * dst_stride1) =
1429 scale
1430 * *src.offset(
1431 ii as isize * src_stride0 + j as isize * src_stride1,
1432 );
1433 }
1434 }
1435 }
1436 }
1437 });
1438 for i in row_full..src_rows {
1439 for j in 0..src_cols {
1440 *dst.offset(j as isize * dst_stride0 + i as isize * dst_stride1) =
1441 scale * *src.offset(i as isize * src_stride0 + j as isize * src_stride1);
1442 }
1443 }
1444 return;
1445 }
1446 }
1447
1448 let mut i = 0;
1449 while i < row_full {
1450 let mut j = 0;
1451 while j < col_full {
1452 transpose_scale_4x4_identity(
1453 dst,
1454 dst_stride0,
1455 dst_stride1,
1456 src,
1457 src_stride0,
1458 src_stride1,
1459 i,
1460 j,
1461 scale,
1462 );
1463 j += TILE;
1464 }
1465 for j in col_full..src_cols {
1466 for ii in i..i + TILE {
1467 *dst.offset(j as isize * dst_stride0 + ii as isize * dst_stride1) =
1468 scale * *src.offset(ii as isize * src_stride0 + j as isize * src_stride1);
1469 }
1470 }
1471 i += TILE;
1472 }
1473 for i in row_full..src_rows {
1474 for j in 0..src_cols {
1475 *dst.offset(j as isize * dst_stride0 + i as isize * dst_stride1) =
1476 scale * *src.offset(i as isize * src_stride0 + j as isize * src_stride1);
1477 }
1478 }
1479}
1480
1481#[inline]
1482unsafe fn try_copy_transpose_scale_2d_identity_tiled<T>(
1483 dest: &mut StridedViewMut<T>,
1484 src: &StridedView<T>,
1485 scale: T,
1486) -> bool
1487where
1488 T: Copy + Mul<Output = T> + MaybeSendSync,
1489{
1490 if src.ndim() != 2 || dest.ndim() != 2 {
1491 return false;
1492 }
1493 let src_dims = src.dims();
1494 if dest.dims() != [src_dims[1], src_dims[0]] {
1495 return false;
1496 }
1497 if src.strides()[0] != 1 || dest.strides()[0] != 1 {
1498 return false;
1499 }
1500
1501 copy_transpose_scale_2d_identity_tiled_raw(
1502 dest.as_mut_ptr(),
1503 dest.strides()[0],
1504 dest.strides()[1],
1505 src.ptr(),
1506 src.strides()[0],
1507 src.strides()[1],
1508 src_dims[0],
1509 src_dims[1],
1510 scale,
1511 );
1512 true
1513}
1514
1515#[inline]
1516unsafe fn copy_transpose_scale_2d_loop<T>(
1517 dst: *mut T,
1518 dst_stride0: isize,
1519 dst_stride1: isize,
1520 src: *const T,
1521 src_stride0: isize,
1522 src_stride1: isize,
1523 src_rows: usize,
1524 src_cols: usize,
1525 scale: T,
1526) where
1527 T: Copy + ElementOpApply + Mul<Output = T> + MaybeSendSync,
1528{
1529 #[cfg(feature = "parallel")]
1530 {
1531 let total = src_rows.saturating_mul(src_cols);
1532 let nthreads = crate::execution_policy::rayon_threads();
1533 if total > MINTHREADLENGTH && nthreads > 1 {
1534 let dst_send = SendPtr(dst);
1535 let src_send = SendPtr(src as *mut T);
1536 if dst_stride0.unsigned_abs() <= dst_stride1.unsigned_abs() {
1537 crate::threading::parallel_for_each(0..src_rows, nthreads, &|rows| {
1538 for i in rows {
1539 let dst = dst_send.as_ptr();
1540 let src = src_send.as_const();
1541 unsafe {
1542 for j in 0..src_cols {
1543 let value = (*src
1544 .offset(i as isize * src_stride0 + j as isize * src_stride1))
1545 .transpose();
1546 *dst.offset(j as isize * dst_stride0 + i as isize * dst_stride1) =
1547 scale * value;
1548 }
1549 }
1550 }
1551 });
1552 } else {
1553 crate::threading::parallel_for_each(0..src_cols, nthreads, &|columns| {
1554 for j in columns {
1555 let dst = dst_send.as_ptr();
1556 let src = src_send.as_const();
1557 unsafe {
1558 for i in 0..src_rows {
1559 let value = (*src
1560 .offset(i as isize * src_stride0 + j as isize * src_stride1))
1561 .transpose();
1562 *dst.offset(j as isize * dst_stride0 + i as isize * dst_stride1) =
1563 scale * value;
1564 }
1565 }
1566 }
1567 });
1568 }
1569 return;
1570 }
1571 }
1572
1573 if dst_stride0.unsigned_abs() <= dst_stride1.unsigned_abs() {
1574 for i in 0..src_rows {
1575 for j in 0..src_cols {
1576 let value =
1577 (*src.offset(i as isize * src_stride0 + j as isize * src_stride1)).transpose();
1578 *dst.offset(j as isize * dst_stride0 + i as isize * dst_stride1) = scale * value;
1579 }
1580 }
1581 } else {
1582 for j in 0..src_cols {
1583 for i in 0..src_rows {
1584 let value =
1585 (*src.offset(i as isize * src_stride0 + j as isize * src_stride1)).transpose();
1586 *dst.offset(j as isize * dst_stride0 + i as isize * dst_stride1) = scale * value;
1587 }
1588 }
1589 }
1590}
1591
1592pub fn copy_transpose_scale_into<T>(
1594 dest: &mut StridedViewMut<T>,
1595 src: &StridedView<T>,
1596 scale: T,
1597) -> Result<()>
1598where
1599 T: Copy + ElementOpApply + Mul<Output = T> + Zero + One + PartialEq + MaybeSendSync + 'static,
1600{
1601 if src.ndim() != 2 || dest.ndim() != 2 {
1602 return Err(StridedError::RankMismatch(src.ndim(), 2));
1603 }
1604
1605 let src_dims = src.dims();
1606 let expected_dims = [src_dims[1], src_dims[0]];
1607 ensure_same_shape(dest.dims(), &expected_dims)?;
1608
1609 if scale == T::zero() {
1610 unsafe {
1611 if same_contiguous_layout(dest.dims(), &[dest.strides()]).is_some() {
1612 fill_contiguous(dest.as_mut_ptr(), total_len(dest.dims()), T::zero());
1613 } else {
1614 fill_2d(
1615 dest.as_mut_ptr(),
1616 dest.dims()[0],
1617 dest.dims()[1],
1618 dest.strides()[0],
1619 dest.strides()[1],
1620 T::zero(),
1621 );
1622 }
1623 }
1624 return Ok(());
1625 }
1626
1627 let transpose_is_identity = element_transpose_is_identity::<T>();
1628
1629 unsafe {
1630 if transpose_is_identity && try_copy_transpose_scale_2d_f64_tiled_typed(dest, src, scale) {
1631 return Ok(());
1632 }
1633 if transpose_is_identity && try_copy_transpose_scale_2d_identity_tiled(dest, src, scale) {
1634 return Ok(());
1635 }
1636 }
1637
1638 if scale == T::one() && transpose_is_identity {
1639 let src_t = src.permute(&[1, 0])?;
1640 #[cfg(feature = "parallel")]
1641 {
1642 return crate::threading::copy_permuted_with_active_policy(dest, &src_t);
1643 }
1644 #[cfg(not(feature = "parallel"))]
1645 return crate::threading::copy_permuted_serial(dest, &src_t);
1646 }
1647
1648 unsafe {
1649 copy_transpose_scale_2d_loop(
1650 dest.as_mut_ptr(),
1651 dest.strides()[0],
1652 dest.strides()[1],
1653 src.ptr(),
1654 src.strides()[0],
1655 src.strides()[1],
1656 src_dims[0],
1657 src_dims[1],
1658 scale,
1659 );
1660 }
1661 Ok(())
1662}
1663
1664#[cfg(test)]
1665#[path = "ops_view/tests/tiled_tests.rs"]
1666mod tiled_tests;