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