1use crate::kernel::{
10 build_plan_fused, build_plan_fused_small, ensure_same_shape, for_each_inner_block_preordered,
11 sequential_contiguous_layout, total_len, SMALL_TENSOR_THRESHOLD,
12};
13use crate::maybe_sync::{MaybeSendSync, MaybeSync};
14use crate::simd;
15use crate::view::{StridedView, StridedViewMut};
16use crate::{Result, StridedError};
17use core::mem::MaybeUninit;
18use std::ops::Mul;
19use strided_view::ElementOp;
20
21#[cfg(feature = "parallel")]
22use crate::fuse::compute_costs;
23#[cfg(feature = "parallel")]
24use crate::threading::{for_each_inner_block_with_offsets, mapreduce_threaded, MINTHREADLENGTH};
25#[cfg(feature = "parallel")]
26use smallvec::SmallVec;
27
28#[cfg(feature = "parallel")]
29type AxisVec<T> = SmallVec<[T; 8]>;
30#[cfg(not(feature = "parallel"))]
31type AxisVec<T> = Vec<T>;
32
33const CONTIGUOUS_RANGE_MIN_LEN: usize = 1 << 15;
34
35#[derive(Clone, Copy, Debug)]
41pub struct ValidatedDestinationLayout(());
42
43#[inline]
44fn validate_destination_layout(
45 dims: &[usize],
46 strides: &[isize],
47) -> Result<ValidatedDestinationLayout> {
48 if crate::layout_check::is_injective_layout(dims, strides) {
49 Ok(ValidatedDestinationLayout(()))
50 } else {
51 Err(StridedError::NonInjectiveOutputLayout)
52 }
53}
54
55#[inline]
56pub fn validate_destination_layout_without_alloc(
70 dims: &[usize],
71 strides: &[isize],
72) -> Result<ValidatedDestinationLayout> {
73 if crate::layout_check::is_injective_layout_without_alloc(dims, strides) {
74 Ok(ValidatedDestinationLayout(()))
75 } else {
76 Err(StridedError::NonInjectiveOutputLayout)
77 }
78}
79
80fn reachable_byte_range(
81 ptr: usize,
82 elem_size: usize,
83 dims: &[usize],
84 strides: &[isize],
85) -> Result<Option<(usize, usize)>> {
86 if elem_size == 0 || dims.contains(&0) {
87 return Ok(None);
88 }
89
90 let mut min_offset = 0isize;
91 let mut max_offset = 0isize;
92 for (&dim, &stride) in dims.iter().zip(strides) {
93 if dim <= 1 {
94 continue;
95 }
96 let extent = isize::try_from(dim - 1).map_err(|_| StridedError::OffsetOverflow)?;
97 let span = stride
98 .checked_mul(extent)
99 .ok_or(StridedError::OffsetOverflow)?;
100 if span < 0 {
101 min_offset = min_offset
102 .checked_add(span)
103 .ok_or(StridedError::OffsetOverflow)?;
104 } else {
105 max_offset = max_offset
106 .checked_add(span)
107 .ok_or(StridedError::OffsetOverflow)?;
108 }
109 }
110
111 let ptr = ptr as i128;
112 let elem_size = elem_size as i128;
113 let start = ptr
114 .checked_add(
115 (min_offset as i128)
116 .checked_mul(elem_size)
117 .ok_or(StridedError::OffsetOverflow)?,
118 )
119 .ok_or(StridedError::OffsetOverflow)?;
120 let end = ptr
121 .checked_add(
122 (max_offset as i128)
123 .checked_mul(elem_size)
124 .and_then(|offset| offset.checked_add(elem_size))
125 .ok_or(StridedError::OffsetOverflow)?,
126 )
127 .ok_or(StridedError::OffsetOverflow)?;
128 if start < 0 || end < 0 || start > usize::MAX as i128 || end > usize::MAX as i128 {
129 return Err(StridedError::OffsetOverflow);
130 }
131 Ok(Some((start as usize, end as usize)))
132}
133
134fn validate_typed_no_overlap<D, A, Op: ElementOp<A>>(
135 dest: &StridedViewMut<MaybeUninit<D>>,
136 input: &StridedView<A, Op>,
137 input_index: usize,
138) -> Result<()> {
139 let dest_range = reachable_byte_range(
140 dest.ptr() as usize,
141 core::mem::size_of::<D>(),
142 dest.dims(),
143 dest.strides(),
144 )?;
145 let input_range = reachable_byte_range(
146 input.ptr() as usize,
147 core::mem::size_of::<A>(),
148 input.dims(),
149 input.strides(),
150 )?;
151 if let (Some((dest_start, dest_end)), Some((input_start, input_end))) =
152 (dest_range, input_range)
153 {
154 if dest_start < input_end && input_start < dest_end {
155 return Err(StridedError::OverlappingInputOutput { input: input_index });
156 }
157 }
158 Ok(())
159}
160
161#[inline(always)]
171unsafe fn inner_loop_map1<D: Copy, A: Copy, Op: ElementOp<A>>(
172 dp: *mut D,
173 ds: isize,
174 sp: *const A,
175 ss: isize,
176 len: usize,
177 f: &impl Fn(A) -> D,
178) {
179 if ds == 1 && ss == 1 {
180 let src = std::slice::from_raw_parts(sp, len);
181 let dst = std::slice::from_raw_parts_mut(dp, len);
182 simd::dispatch_if_large(len, || {
183 for (d, s) in dst.iter_mut().zip(src.iter()) {
184 *d = f(Op::apply(*s));
185 }
186 });
187 } else {
188 let mut dp = dp;
189 let mut sp = sp;
190 for _ in 0..len {
191 *dp = f(Op::apply(*sp));
192 dp = dp.offset(ds);
193 sp = sp.offset(ss);
194 }
195 }
196}
197
198#[inline(always)]
200unsafe fn inner_loop_map2<D: Copy, A: Copy, B: Copy, OpA: ElementOp<A>, OpB: ElementOp<B>>(
201 dp: *mut D,
202 ds: isize,
203 ap: *const A,
204 a_s: isize,
205 bp: *const B,
206 b_s: isize,
207 len: usize,
208 f: &impl Fn(A, B) -> D,
209) {
210 if ds == 1 && a_s == 1 && b_s == 1 {
211 let src_a = std::slice::from_raw_parts(ap, len);
212 let src_b = std::slice::from_raw_parts(bp, len);
213 let dst = std::slice::from_raw_parts_mut(dp, len);
214 simd::dispatch_if_large(len, || {
215 for i in 0..len {
216 dst[i] = f(OpA::apply(src_a[i]), OpB::apply(src_b[i]));
217 }
218 });
219 } else if ds == 1 && a_s == 1 && b_s == 0 {
220 let src_a = std::slice::from_raw_parts(ap, len);
221 let b = OpB::apply(*bp);
222 let dst = std::slice::from_raw_parts_mut(dp, len);
223 simd::dispatch_if_large(len, || {
224 for i in 0..len {
225 dst[i] = f(OpA::apply(src_a[i]), b);
226 }
227 });
228 } else if ds == 1 && a_s == 0 && b_s == 1 {
229 let a = OpA::apply(*ap);
230 let src_b = std::slice::from_raw_parts(bp, len);
231 let dst = std::slice::from_raw_parts_mut(dp, len);
232 simd::dispatch_if_large(len, || {
233 for i in 0..len {
234 dst[i] = f(a, OpB::apply(src_b[i]));
235 }
236 });
237 } else if ds == 1 && a_s == 0 && b_s == 0 {
238 let a = OpA::apply(*ap);
239 let b = OpB::apply(*bp);
240 let dst = std::slice::from_raw_parts_mut(dp, len);
241 simd::dispatch_if_large(len, || {
242 for d in dst.iter_mut() {
243 *d = f(a, b);
244 }
245 });
246 } else if ds == 1 && b_s == 0 {
247 let b = OpB::apply(*bp);
248 let dst = std::slice::from_raw_parts_mut(dp, len);
249 let mut ap = ap;
250 simd::dispatch_if_large(len, || {
251 for d in dst.iter_mut() {
252 *d = f(OpA::apply(*ap), b);
253 ap = ap.offset(a_s);
254 }
255 });
256 } else if ds == 1 && a_s == 0 {
257 let a = OpA::apply(*ap);
258 let dst = std::slice::from_raw_parts_mut(dp, len);
259 let mut bp = bp;
260 simd::dispatch_if_large(len, || {
261 for d in dst.iter_mut() {
262 *d = f(a, OpB::apply(*bp));
263 bp = bp.offset(b_s);
264 }
265 });
266 } else {
267 let mut dp = dp;
268 let mut ap = ap;
269 let mut bp = bp;
270 for _ in 0..len {
271 *dp = f(OpA::apply(*ap), OpB::apply(*bp));
272 dp = dp.offset(ds);
273 ap = ap.offset(a_s);
274 bp = bp.offset(b_s);
275 }
276 }
277}
278
279trait MulOutput<D: Copy + 'static>: Copy + MaybeSendSync + 'static {
281 type Slot: Copy + MaybeSendSync + 'static;
282
283 unsafe fn write(dst: *mut Self::Slot, value: D);
284
285 unsafe fn try_contiguous<A: 'static, B: 'static>(
286 dst: *mut Self::Slot,
287 len: usize,
288 a: &[A],
289 b: &[B],
290 ) -> bool;
291}
292
293#[derive(Clone, Copy)]
294struct InitializedOutput;
295
296impl<D: Copy + MaybeSendSync + 'static> MulOutput<D> for InitializedOutput {
297 type Slot = D;
298
299 #[inline(always)]
300 unsafe fn write(dst: *mut D, value: D) {
301 unsafe { dst.write(value) };
302 }
303
304 #[inline(always)]
305 unsafe fn try_contiguous<A: 'static, B: 'static>(
306 dst: *mut D,
307 len: usize,
308 a: &[A],
309 b: &[B],
310 ) -> bool {
311 unsafe { simd::try_mul_contiguous_ptr(dst, len, a, b) }
312 }
313}
314
315#[derive(Clone, Copy)]
316struct UninitializedOutput;
317
318impl<D: Copy + MaybeSendSync + 'static> MulOutput<D> for UninitializedOutput {
319 type Slot = MaybeUninit<D>;
320
321 #[inline(always)]
322 unsafe fn write(dst: *mut MaybeUninit<D>, value: D) {
323 unsafe { dst.write(MaybeUninit::new(value)) };
324 }
325
326 #[inline(always)]
327 unsafe fn try_contiguous<A: 'static, B: 'static>(
328 dst: *mut MaybeUninit<D>,
329 len: usize,
330 a: &[A],
331 b: &[B],
332 ) -> bool {
333 unsafe { simd::try_mul_contiguous_ptr(dst.cast::<D>(), len, a, b) }
336 }
337}
338
339#[inline(always)]
340fn multiply_value<A, B, D>(lhs: A, rhs: B) -> D
341where
342 A: Copy + Mul<B, Output = D> + 'static,
343 B: Copy + 'static,
344 D: Copy + 'static,
345{
346 use core::any::TypeId;
347
348 if TypeId::of::<A>() == TypeId::of::<i32>()
349 && TypeId::of::<B>() == TypeId::of::<i32>()
350 && TypeId::of::<D>() == TypeId::of::<i32>()
351 {
352 let lhs = unsafe { *(&lhs as *const A).cast::<i32>() };
354 let rhs = unsafe { *(&rhs as *const B).cast::<i32>() };
355 let value = lhs.wrapping_mul(rhs);
356 return unsafe { core::mem::transmute_copy(&value) };
357 }
358 if TypeId::of::<A>() == TypeId::of::<i64>()
359 && TypeId::of::<B>() == TypeId::of::<i64>()
360 && TypeId::of::<D>() == TypeId::of::<i64>()
361 {
362 let lhs = unsafe { *(&lhs as *const A).cast::<i64>() };
364 let rhs = unsafe { *(&rhs as *const B).cast::<i64>() };
365 let value = lhs.wrapping_mul(rhs);
366 return unsafe { core::mem::transmute_copy(&value) };
367 }
368 lhs * rhs
369}
370
371#[inline(always)]
372unsafe fn inner_loop_mul2<
373 O: MulOutput<D>,
374 D: Copy + 'static,
375 A: Copy + Mul<B, Output = D> + 'static,
376 B: Copy + 'static,
377>(
378 dp: *mut O::Slot,
379 ds: isize,
380 ap: *const A,
381 a_s: isize,
382 bp: *const B,
383 b_s: isize,
384 len: usize,
385) {
386 if ds == 1 && a_s == 1 && b_s == 1 {
387 let src_a = std::slice::from_raw_parts(ap, len);
388 let src_b = std::slice::from_raw_parts(bp, len);
389 if len >= 64 && O::try_contiguous(dp, len, src_a, src_b) {
390 return;
391 }
392 for i in 0..len {
393 O::write(dp.add(i), multiply_value(src_a[i], src_b[i]));
394 }
395 } else if ds == 1 && a_s == 1 && b_s == 0 {
396 let src_a = std::slice::from_raw_parts(ap, len);
397 let b = *bp;
398 for i in 0..len {
399 O::write(dp.add(i), multiply_value(src_a[i], b));
400 }
401 } else if ds == 1 && a_s == 0 && b_s == 1 {
402 let a = *ap;
403 let src_b = std::slice::from_raw_parts(bp, len);
404 for i in 0..len {
405 O::write(dp.add(i), multiply_value(a, src_b[i]));
406 }
407 } else if ds == 1 && a_s == 0 && b_s == 0 {
408 let a = *ap;
409 let b = *bp;
410 for i in 0..len {
411 O::write(dp.add(i), multiply_value(a, b));
412 }
413 } else if ds == 1 && b_s == 0 {
414 let b = *bp;
415 let mut ap = ap;
416 for i in 0..len {
417 O::write(dp.add(i), multiply_value(*ap, b));
418 ap = ap.offset(a_s);
419 }
420 } else if ds == 1 && a_s == 0 {
421 let a = *ap;
422 let mut bp = bp;
423 for i in 0..len {
424 O::write(dp.add(i), multiply_value(a, *bp));
425 bp = bp.offset(b_s);
426 }
427 } else {
428 let mut dp = dp;
429 let mut ap = ap;
430 let mut bp = bp;
431 for _ in 0..len {
432 O::write(dp, multiply_value(*ap, *bp));
433 dp = dp.offset(ds);
434 ap = ap.offset(a_s);
435 bp = bp.offset(b_s);
436 }
437 }
438}
439
440#[derive(Clone, Debug, Eq, PartialEq)]
441struct ContiguousMulRangePlan {
442 axis_order: AxisVec<usize>,
443 inner_len: usize,
444 inner_axis_count: usize,
445 row_len: usize,
446 outer_axis_start: usize,
447 fast_axis: usize,
448 a_fast_stride: isize,
449 b_fast_stride: isize,
450 a_row_stride: isize,
451 b_row_stride: isize,
452}
453
454#[cfg(feature = "parallel")]
455#[derive(Clone, Copy, Debug, Eq, PartialEq)]
456enum TransposedScalarTileKind {
457 RhsScalar,
458 LhsScalar,
459}
460
461#[cfg(feature = "parallel")]
462fn transposed_scalar_tile_kind(plan: &ContiguousMulRangePlan) -> Option<TransposedScalarTileKind> {
463 let row_len = isize::try_from(plan.row_len).ok()?;
464 if plan.a_fast_stride == row_len
465 && plan.a_row_stride == 1
466 && plan.b_fast_stride == 0
467 && plan.b_row_stride == 0
468 {
469 return Some(TransposedScalarTileKind::RhsScalar);
470 }
471
472 if plan.b_fast_stride == row_len
473 && plan.b_row_stride == 1
474 && plan.a_fast_stride == 0
475 && plan.a_row_stride == 0
476 {
477 return Some(TransposedScalarTileKind::LhsScalar);
478 }
479
480 None
481}
482
483fn compact_axis_order(dims: &[usize], strides: &[isize]) -> Option<AxisVec<usize>> {
484 if dims.len() != strides.len() {
485 return None;
486 }
487
488 let mut active = AxisVec::<usize>::new();
489 let mut inactive = AxisVec::<usize>::new();
490 for (axis, (&dim, &stride)) in dims.iter().zip(strides.iter()).enumerate() {
491 if stride < 0 {
492 return None;
493 }
494 if dim > 1 {
495 active.push(axis);
496 } else {
497 inactive.push(axis);
498 }
499 }
500
501 active.sort_by(|&lhs, &rhs| strides[lhs].cmp(&strides[rhs]).then_with(|| lhs.cmp(&rhs)));
502
503 let mut expected = 1isize;
504 for &axis in &active {
505 if strides[axis] != expected {
506 return None;
507 }
508 expected = expected.saturating_mul(dims[axis] as isize);
509 }
510
511 active.extend(inactive);
512 Some(active)
513}
514
515fn can_fuse_contiguous_range_axis(dim: usize, prev_stride: isize, next_stride: isize) -> bool {
516 dim <= 1 || (prev_stride == 0 && next_stride == 0) || next_stride == prev_stride * dim as isize
517}
518
519fn contiguous_mul_range_plan(
520 dims: &[usize],
521 dst_strides: &[isize],
522 a_strides: &[isize],
523 b_strides: &[isize],
524) -> Option<ContiguousMulRangePlan> {
525 let axis_order = compact_axis_order(dims, dst_strides)?;
526 if dims.is_empty() {
527 return Some(ContiguousMulRangePlan {
528 axis_order,
529 inner_len: 1,
530 inner_axis_count: 0,
531 row_len: 1,
532 outer_axis_start: 0,
533 fast_axis: 0,
534 a_fast_stride: 0,
535 b_fast_stride: 0,
536 a_row_stride: 0,
537 b_row_stride: 0,
538 });
539 }
540
541 let first_pos = axis_order
542 .iter()
543 .position(|&axis| dims[axis] > 1)
544 .unwrap_or(0);
545 let first_axis = axis_order[first_pos];
546 let mut inner_len = dims[first_axis].max(1);
547 let mut inner_axis_count = first_pos + 1;
548 let mut prev_axis = first_axis;
549
550 for &axis in axis_order.iter().skip(first_pos + 1) {
551 if can_fuse_contiguous_range_axis(
552 dims[prev_axis],
553 dst_strides[prev_axis],
554 dst_strides[axis],
555 ) && can_fuse_contiguous_range_axis(
556 dims[prev_axis],
557 a_strides[prev_axis],
558 a_strides[axis],
559 ) && can_fuse_contiguous_range_axis(
560 dims[prev_axis],
561 b_strides[prev_axis],
562 b_strides[axis],
563 ) {
564 inner_len = inner_len.checked_mul(dims[axis].max(1))?;
565 inner_axis_count += 1;
566 prev_axis = axis;
567 } else {
568 break;
569 }
570 }
571
572 let row = axis_order
573 .iter()
574 .enumerate()
575 .skip(inner_axis_count)
576 .find(|&(_, &axis)| dims[axis] > 1);
577 let (row_len, outer_axis_start, a_row_stride, b_row_stride) =
578 if let Some((row_pos, &row_axis)) = row {
579 (
580 dims[row_axis],
581 row_pos + 1,
582 a_strides[row_axis],
583 b_strides[row_axis],
584 )
585 } else {
586 (1, axis_order.len(), 0, 0)
587 };
588
589 Some(ContiguousMulRangePlan {
590 axis_order,
591 inner_len,
592 inner_axis_count,
593 row_len,
594 outer_axis_start,
595 fast_axis: first_axis,
596 a_fast_stride: a_strides[first_axis],
597 b_fast_stride: b_strides[first_axis],
598 a_row_stride,
599 b_row_stride,
600 })
601}
602
603struct ContiguousMulOuterCursor<'a> {
604 dims: &'a [usize],
605 a_strides: &'a [isize],
606 b_strides: &'a [isize],
607 axes: AxisVec<usize>,
608 coords: AxisVec<usize>,
609 a_offset: isize,
610 b_offset: isize,
611}
612
613impl<'a> ContiguousMulOuterCursor<'a> {
614 fn new(
615 dims: &'a [usize],
616 a_strides: &'a [isize],
617 b_strides: &'a [isize],
618 plan: &ContiguousMulRangePlan,
619 outer_group: usize,
620 ) -> Self {
621 let axes: AxisVec<usize> = plan
622 .axis_order
623 .iter()
624 .skip(plan.outer_axis_start)
625 .copied()
626 .collect();
627 let mut coords = AxisVec::<usize>::with_capacity(axes.len());
628 let mut rem = outer_group;
629 let mut a_offset = 0isize;
630 let mut b_offset = 0isize;
631
632 for &axis in &axes {
633 let dim = dims[axis].max(1);
634 let coord = rem % dim;
635 rem /= dim;
636 coords.push(coord);
637 a_offset += coord as isize * a_strides[axis];
638 b_offset += coord as isize * b_strides[axis];
639 }
640
641 Self {
642 dims,
643 a_strides,
644 b_strides,
645 axes,
646 coords,
647 a_offset,
648 b_offset,
649 }
650 }
651
652 fn advance(&mut self) {
653 for (i, &axis) in self.axes.iter().enumerate() {
654 let dim = self.dims[axis].max(1);
655 if dim <= 1 {
656 continue;
657 }
658
659 self.coords[i] += 1;
660 self.a_offset += self.a_strides[axis];
661 self.b_offset += self.b_strides[axis];
662
663 if self.coords[i] < dim {
664 break;
665 }
666
667 self.coords[i] = 0;
668 self.a_offset -= dim as isize * self.a_strides[axis];
669 self.b_offset -= dim as isize * self.b_strides[axis];
670 }
671 }
672}
673
674#[inline(always)]
675unsafe fn run_contiguous_mul_row_block<
676 O: MulOutput<D>,
677 D: Copy + 'static,
678 A: Copy + Mul<B, Output = D> + 'static,
679 B: Copy + 'static,
680>(
681 dst_ptr: *mut O::Slot,
682 a_ptr: *const A,
683 b_ptr: *const B,
684 plan: &ContiguousMulRangePlan,
685 base_index: usize,
686 total: usize,
687 base_a_offset: isize,
688 base_b_offset: isize,
689) {
690 let inner_len = plan.inner_len.max(1);
691 let row_len = plan.row_len.max(1);
692 #[cfg(feature = "parallel")]
693 let block_len = inner_len.saturating_mul(row_len);
694
695 #[cfg(feature = "parallel")]
696 if total.saturating_sub(base_index) >= block_len {
697 match transposed_scalar_tile_kind(plan) {
698 Some(TransposedScalarTileKind::RhsScalar) => {
699 if simd::try_mul_transposed_scalar_rhs_2d::<D, A, B>(
700 dst_ptr.add(base_index).cast::<D>(),
701 a_ptr.offset(base_a_offset),
702 b_ptr.offset(base_b_offset),
703 inner_len,
704 row_len,
705 plan.a_fast_stride,
706 plan.a_row_stride,
707 ) {
708 return;
709 }
710 }
711 Some(TransposedScalarTileKind::LhsScalar) => {
712 if simd::try_mul_transposed_scalar_lhs_2d::<D, A, B>(
713 dst_ptr.add(base_index).cast::<D>(),
714 a_ptr.offset(base_a_offset),
715 b_ptr.offset(base_b_offset),
716 inner_len,
717 row_len,
718 plan.b_fast_stride,
719 plan.b_row_stride,
720 ) {
721 return;
722 }
723 }
724 None => {}
725 }
726 }
727
728 let mut index = base_index;
729 let mut a_offset = base_a_offset;
730 let mut b_offset = base_b_offset;
731
732 for _ in 0..row_len {
733 if index >= total {
734 break;
735 }
736 let len = inner_len.min(total - index);
737 inner_loop_mul2::<O, D, A, B>(
738 dst_ptr.add(index),
739 1,
740 a_ptr.offset(a_offset),
741 plan.a_fast_stride,
742 b_ptr.offset(b_offset),
743 plan.b_fast_stride,
744 len,
745 );
746 index += inner_len;
747 a_offset += plan.a_row_stride;
748 b_offset += plan.b_row_stride;
749 }
750}
751
752#[cfg(feature = "parallel")]
753fn strided_offset_for_contiguous_linear_index(
754 dims: &[usize],
755 strides: &[isize],
756 axis_order: &[usize],
757 mut index: usize,
758) -> isize {
759 let mut offset = 0isize;
760 for &axis in axis_order {
761 let dim = dims[axis];
762 if dim == 0 {
763 return 0;
764 }
765 let coord = index % dim;
766 index /= dim;
767 offset += coord as isize * strides[axis];
768 }
769 offset
770}
771
772fn try_contiguous_range_mul<
773 O: MulOutput<D>,
774 D: Copy + MaybeSendSync + 'static,
775 A: Copy + MaybeSendSync + Mul<B, Output = D> + 'static,
776 B: Copy + MaybeSendSync + 'static,
777>(
778 dst_ptr: *mut O::Slot,
779 dims: &[usize],
780 dst_strides: &[isize],
781 a_ptr: *const A,
782 a_strides: &[isize],
783 b_ptr: *const B,
784 b_strides: &[isize],
785) -> bool {
786 let total = total_len(dims);
787 if total == 0 {
788 return true;
789 }
790 if total <= CONTIGUOUS_RANGE_MIN_LEN {
791 return false;
792 }
793
794 let Some(plan) = contiguous_mul_range_plan(dims, dst_strides, a_strides, b_strides) else {
795 return false;
796 };
797
798 let inner_len = plan.inner_len.max(1);
799 let row_len = plan.row_len.max(1);
800 let block_len = inner_len.saturating_mul(row_len).max(1);
801 let outer_groups = total.div_ceil(block_len);
802
803 #[cfg(feature = "parallel")]
804 {
805 let nthreads = crate::execution_policy::rayon_threads();
806 if nthreads > 1 {
807 use crate::threading::{parallel_for_each, SendPtr};
808
809 let dst = SendPtr(dst_ptr);
810 let a = SendPtr(a_ptr as *mut A);
811 let b = SendPtr(b_ptr as *mut B);
812
813 if outer_groups < nthreads {
814 let chunk_len = total.div_ceil(nthreads);
815 let nchunks = total.div_ceil(chunk_len);
816
817 parallel_for_each(0..nchunks, nthreads, &|chunks| {
818 for chunk in chunks {
819 let start = chunk * chunk_len;
820 let end = (start + chunk_len).min(total);
821 let mut index = start;
822
823 while index < end {
824 let in_inner = index % inner_len;
825 let len = (inner_len - in_inner).min(end - index);
826 let a_offset = strided_offset_for_contiguous_linear_index(
827 dims,
828 a_strides,
829 &plan.axis_order,
830 index,
831 );
832 let b_offset = strided_offset_for_contiguous_linear_index(
833 dims,
834 b_strides,
835 &plan.axis_order,
836 index,
837 );
838
839 unsafe {
840 inner_loop_mul2::<O, D, A, B>(
841 dst.as_ptr().add(index),
842 1,
843 a.as_const().offset(a_offset),
844 plan.a_fast_stride,
845 b.as_const().offset(b_offset),
846 plan.b_fast_stride,
847 len,
848 );
849 }
850 index += len;
851 }
852 }
853 });
854
855 return true;
856 }
857
858 let groups_per_chunk = outer_groups.div_ceil(nthreads);
859 let nchunks = outer_groups.div_ceil(groups_per_chunk);
860
861 parallel_for_each(0..nchunks, nthreads, &|chunks| {
862 for chunk in chunks {
863 let group_start = chunk * groups_per_chunk;
864 let group_end = (group_start + groups_per_chunk).min(outer_groups);
865 let mut cursor = ContiguousMulOuterCursor::new(
866 dims,
867 a_strides,
868 b_strides,
869 &plan,
870 group_start,
871 );
872
873 for group in group_start..group_end {
874 let index = group * block_len;
875 unsafe {
876 run_contiguous_mul_row_block::<O, D, A, B>(
877 dst.as_ptr(),
878 a.as_const(),
879 b.as_const(),
880 &plan,
881 index,
882 total,
883 cursor.a_offset,
884 cursor.b_offset,
885 );
886 }
887 cursor.advance();
888 }
889 }
890 });
891
892 true
893 } else {
894 run_contiguous_range_mul_single_thread::<O, D, A, B>(
895 dst_ptr,
896 dims,
897 a_ptr,
898 a_strides,
899 b_ptr,
900 b_strides,
901 &plan,
902 total,
903 block_len,
904 outer_groups,
905 )
906 }
907 }
908
909 #[cfg(not(feature = "parallel"))]
910 {
911 run_contiguous_range_mul_single_thread::<O, D, A, B>(
912 dst_ptr,
913 dims,
914 a_ptr,
915 a_strides,
916 b_ptr,
917 b_strides,
918 &plan,
919 total,
920 block_len,
921 outer_groups,
922 )
923 }
924}
925
926fn run_contiguous_range_mul_single_thread<
927 O: MulOutput<D>,
928 D: Copy + MaybeSendSync + 'static,
929 A: Copy + MaybeSendSync + Mul<B, Output = D> + 'static,
930 B: Copy + MaybeSendSync + 'static,
931>(
932 dst_ptr: *mut O::Slot,
933 dims: &[usize],
934 a_ptr: *const A,
935 a_strides: &[isize],
936 b_ptr: *const B,
937 b_strides: &[isize],
938 plan: &ContiguousMulRangePlan,
939 total: usize,
940 block_len: usize,
941 outer_groups: usize,
942) -> bool {
943 let mut cursor = ContiguousMulOuterCursor::new(dims, a_strides, b_strides, plan, 0);
944 for group in 0..outer_groups {
945 let index = group * block_len;
946 unsafe {
947 run_contiguous_mul_row_block::<O, D, A, B>(
948 dst_ptr,
949 a_ptr,
950 b_ptr,
951 plan,
952 index,
953 total,
954 cursor.a_offset,
955 cursor.b_offset,
956 );
957 }
958 cursor.advance();
959 }
960 true
961}
962
963#[inline(always)]
965unsafe fn inner_loop_map3<
966 D: Copy,
967 A: Copy,
968 B: Copy,
969 C: Copy,
970 OpA: ElementOp<A>,
971 OpB: ElementOp<B>,
972 OpC: ElementOp<C>,
973>(
974 dp: *mut D,
975 ds: isize,
976 ap: *const A,
977 a_s: isize,
978 bp: *const B,
979 b_s: isize,
980 cp: *const C,
981 c_s: isize,
982 len: usize,
983 f: &impl Fn(A, B, C) -> D,
984) {
985 if ds == 1 && a_s == 1 && b_s == 1 && c_s == 1 {
986 let src_a = std::slice::from_raw_parts(ap, len);
987 let src_b = std::slice::from_raw_parts(bp, len);
988 let src_c = std::slice::from_raw_parts(cp, len);
989 let dst = std::slice::from_raw_parts_mut(dp, len);
990 simd::dispatch_if_large(len, || {
991 for i in 0..len {
992 dst[i] = f(
993 OpA::apply(src_a[i]),
994 OpB::apply(src_b[i]),
995 OpC::apply(src_c[i]),
996 );
997 }
998 });
999 } else {
1000 let mut dp = dp;
1001 let mut ap = ap;
1002 let mut bp = bp;
1003 let mut cp = cp;
1004 for _ in 0..len {
1005 *dp = f(OpA::apply(*ap), OpB::apply(*bp), OpC::apply(*cp));
1006 dp = dp.offset(ds);
1007 ap = ap.offset(a_s);
1008 bp = bp.offset(b_s);
1009 cp = cp.offset(c_s);
1010 }
1011 }
1012}
1013
1014#[inline(always)]
1016unsafe fn inner_loop_map4<
1017 D: Copy,
1018 A: Copy,
1019 B: Copy,
1020 C: Copy,
1021 E: Copy,
1022 OpA: ElementOp<A>,
1023 OpB: ElementOp<B>,
1024 OpC: ElementOp<C>,
1025 OpE: ElementOp<E>,
1026>(
1027 dp: *mut D,
1028 ds: isize,
1029 ap: *const A,
1030 a_s: isize,
1031 bp: *const B,
1032 b_s: isize,
1033 cp: *const C,
1034 c_s: isize,
1035 ep: *const E,
1036 e_s: isize,
1037 len: usize,
1038 f: &impl Fn(A, B, C, E) -> D,
1039) {
1040 if ds == 1 && a_s == 1 && b_s == 1 && c_s == 1 && e_s == 1 {
1041 let src_a = std::slice::from_raw_parts(ap, len);
1042 let src_b = std::slice::from_raw_parts(bp, len);
1043 let src_c = std::slice::from_raw_parts(cp, len);
1044 let src_e = std::slice::from_raw_parts(ep, len);
1045 let dst = std::slice::from_raw_parts_mut(dp, len);
1046 simd::dispatch_if_large(len, || {
1047 for i in 0..len {
1048 dst[i] = f(
1049 OpA::apply(src_a[i]),
1050 OpB::apply(src_b[i]),
1051 OpC::apply(src_c[i]),
1052 OpE::apply(src_e[i]),
1053 );
1054 }
1055 });
1056 } else {
1057 let mut dp = dp;
1058 let mut ap = ap;
1059 let mut bp = bp;
1060 let mut cp = cp;
1061 let mut ep = ep;
1062 for _ in 0..len {
1063 *dp = f(
1064 OpA::apply(*ap),
1065 OpB::apply(*bp),
1066 OpC::apply(*cp),
1067 OpE::apply(*ep),
1068 );
1069 dp = dp.offset(ds);
1070 ap = ap.offset(a_s);
1071 bp = bp.offset(b_s);
1072 cp = cp.offset(c_s);
1073 ep = ep.offset(e_s);
1074 }
1075 }
1076}
1077
1078pub fn map_into<D: Copy + MaybeSendSync, A: Copy + MaybeSendSync, Op: ElementOp<A>>(
1083 dest: &mut StridedViewMut<D>,
1084 src: &StridedView<A, Op>,
1085 f: impl Fn(A) -> D + MaybeSync,
1086) -> Result<()> {
1087 map_parts_into::<D, A, Op>(
1088 dest.as_mut_ptr(),
1089 dest.dims(),
1090 dest.strides(),
1091 src.ptr(),
1092 src.dims(),
1093 src.strides(),
1094 f,
1095 )
1096}
1097
1098pub(crate) fn map_into_validated<
1099 D: Copy + MaybeSendSync,
1100 A: Copy + MaybeSendSync,
1101 Op: ElementOp<A>,
1102>(
1103 dest: &mut StridedViewMut<D>,
1104 src: &StridedView<A, Op>,
1105 f: impl Fn(A) -> D + MaybeSync,
1106 validated: ValidatedDestinationLayout,
1107) -> Result<()> {
1108 ensure_same_shape(dest.dims(), src.dims())?;
1109 map_parts_into_validated::<D, A, Op>(
1110 dest.as_mut_ptr(),
1111 dest.dims(),
1112 dest.strides(),
1113 src.ptr(),
1114 src.strides(),
1115 f,
1116 validated,
1117 )
1118}
1119
1120pub(crate) fn map_raw_into<D: Copy + MaybeSendSync, A: Copy + MaybeSendSync, Op: ElementOp<A>>(
1121 dest: &mut crate::RawStridedMut<'_, D>,
1122 src: &crate::RawStridedRef<'_, A>,
1123 f: impl Fn(A) -> D + MaybeSync,
1124) -> Result<()> {
1125 map_parts_into::<D, A, Op>(
1126 dest.as_mut_ptr(),
1127 dest.dims(),
1128 dest.strides(),
1129 src.ptr(),
1130 src.dims(),
1131 src.strides(),
1132 f,
1133 )
1134}
1135
1136pub(crate) fn map_raw_into_validated<
1137 D: Copy + MaybeSendSync,
1138 A: Copy + MaybeSendSync,
1139 Op: ElementOp<A>,
1140>(
1141 dest: &mut crate::RawStridedMut<'_, D>,
1142 src: &crate::RawStridedRef<'_, A>,
1143 f: impl Fn(A) -> D + MaybeSync,
1144 validated: ValidatedDestinationLayout,
1145) -> Result<()> {
1146 ensure_same_shape(dest.dims(), src.dims())?;
1147 map_parts_into_validated::<D, A, Op>(
1148 dest.as_mut_ptr(),
1149 dest.dims(),
1150 dest.strides(),
1151 src.ptr(),
1152 src.strides(),
1153 f,
1154 validated,
1155 )
1156}
1157
1158#[allow(clippy::too_many_arguments)]
1159fn map_parts_into<D: Copy + MaybeSendSync, A: Copy + MaybeSendSync, Op: ElementOp<A>>(
1160 dst_ptr: *mut D,
1161 dst_dims: &[usize],
1162 dst_strides: &[isize],
1163 src_ptr: *const A,
1164 src_dims: &[usize],
1165 src_strides: &[isize],
1166 f: impl Fn(A) -> D + MaybeSync,
1167) -> Result<()> {
1168 ensure_same_shape(dst_dims, src_dims)?;
1169 let validated = validate_destination_layout(dst_dims, dst_strides)?;
1170 map_parts_into_validated::<D, A, Op>(
1171 dst_ptr,
1172 dst_dims,
1173 dst_strides,
1174 src_ptr,
1175 src_strides,
1176 f,
1177 validated,
1178 )
1179}
1180
1181#[allow(clippy::too_many_arguments)]
1182fn map_parts_into_validated<D: Copy + MaybeSendSync, A: Copy + MaybeSendSync, Op: ElementOp<A>>(
1183 dst_ptr: *mut D,
1184 dst_dims: &[usize],
1185 dst_strides: &[isize],
1186 src_ptr: *const A,
1187 src_strides: &[isize],
1188 f: impl Fn(A) -> D + MaybeSync,
1189 _validated: ValidatedDestinationLayout,
1190) -> Result<()> {
1191 if sequential_contiguous_layout(dst_dims, &[dst_strides, src_strides]).is_some() {
1192 let len = total_len(dst_dims);
1193 let dst = unsafe { std::slice::from_raw_parts_mut(dst_ptr, len) };
1194 let src = unsafe { std::slice::from_raw_parts(src_ptr, len) };
1195 simd::dispatch_if_large(len, || {
1196 for i in 0..len {
1197 dst[i] = f(Op::apply(src[i]));
1198 }
1199 });
1200 return Ok(());
1201 }
1202
1203 let strides_list: [&[isize]; 2] = [dst_strides, src_strides];
1204 let elem_size = std::mem::size_of::<D>().max(std::mem::size_of::<A>());
1205 let total = total_len(dst_dims);
1206
1207 let (fused_dims, ordered_strides, plan) = if total <= SMALL_TENSOR_THRESHOLD {
1209 build_plan_fused_small(dst_dims, &strides_list)
1210 } else {
1211 build_plan_fused(dst_dims, &strides_list, Some(0), elem_size)
1212 };
1213
1214 #[cfg(feature = "parallel")]
1215 {
1216 let total: usize = fused_dims.iter().product();
1217 let nthreads = crate::execution_policy::rayon_threads();
1218 if total > MINTHREADLENGTH && nthreads > 1 {
1219 use crate::threading::SendPtr;
1220 let dst_send = SendPtr(dst_ptr);
1221 let src_send = SendPtr(src_ptr as *mut A);
1222
1223 let costs = compute_costs(&ordered_strides);
1224 let initial_offsets = vec![0isize; strides_list.len()];
1225 return mapreduce_threaded(
1226 &fused_dims,
1227 &plan.block,
1228 &ordered_strides,
1229 &initial_offsets,
1230 &costs,
1231 nthreads,
1232 0,
1233 1,
1234 &|dims, blocks, strides_list, offsets| {
1235 for_each_inner_block_with_offsets(
1236 dims,
1237 blocks,
1238 strides_list,
1239 offsets,
1240 |offsets, len, strides| {
1241 let dp = unsafe { dst_send.as_ptr().offset(offsets[0]) };
1242 let sp = unsafe { src_send.as_const().offset(offsets[1]) };
1243 unsafe {
1244 inner_loop_map1::<D, A, Op>(dp, strides[0], sp, strides[1], len, &f)
1245 };
1246 Ok(())
1247 },
1248 )
1249 },
1250 );
1251 }
1252 }
1253
1254 let initial_offsets = vec![0isize; ordered_strides.len()];
1255 for_each_inner_block_preordered(
1256 &fused_dims,
1257 &plan.block,
1258 &ordered_strides,
1259 &initial_offsets,
1260 |offsets, len, strides| {
1261 let dp = unsafe { dst_ptr.offset(offsets[0]) };
1262 let sp = unsafe { src_ptr.offset(offsets[1]) };
1263 unsafe { inner_loop_map1::<D, A, Op>(dp, strides[0], sp, strides[1], len, &f) };
1264 Ok(())
1265 },
1266 )
1267}
1268
1269pub fn zip_map2_into<
1274 D: Copy + MaybeSendSync,
1275 A: Copy + MaybeSendSync,
1276 B: Copy + MaybeSendSync,
1277 OpA: ElementOp<A>,
1278 OpB: ElementOp<B>,
1279>(
1280 dest: &mut StridedViewMut<D>,
1281 a: &StridedView<A, OpA>,
1282 b: &StridedView<B, OpB>,
1283 f: impl Fn(A, B) -> D + MaybeSync,
1284) -> Result<()> {
1285 zip_map2_parts_into::<D, A, B, OpA, OpB>(
1286 dest.as_mut_ptr(),
1287 dest.dims(),
1288 dest.strides(),
1289 a.ptr(),
1290 a.dims(),
1291 a.strides(),
1292 b.ptr(),
1293 b.dims(),
1294 b.strides(),
1295 f,
1296 )
1297}
1298
1299pub(crate) fn zip_map2_into_validated<
1300 D: Copy + MaybeSendSync,
1301 A: Copy + MaybeSendSync,
1302 B: Copy + MaybeSendSync,
1303 OpA: ElementOp<A>,
1304 OpB: ElementOp<B>,
1305>(
1306 dest: &mut StridedViewMut<D>,
1307 a: &StridedView<A, OpA>,
1308 b: &StridedView<B, OpB>,
1309 f: impl Fn(A, B) -> D + MaybeSync,
1310 validated: ValidatedDestinationLayout,
1311) -> Result<()> {
1312 zip_map2_parts_into_validated::<D, A, B, OpA, OpB>(
1313 dest.as_mut_ptr(),
1314 dest.dims(),
1315 dest.strides(),
1316 a.ptr(),
1317 a.strides(),
1318 b.ptr(),
1319 b.strides(),
1320 f,
1321 validated,
1322 )
1323}
1324
1325#[non_exhaustive]
1327#[derive(Clone, Copy, Debug, Eq, PartialEq)]
1328pub enum CompareOp {
1329 Eq,
1330 Lt,
1331 Le,
1332 Gt,
1333 Ge,
1334}
1335
1336pub fn compare_into<T, OpA, OpB>(
1347 dest: &mut StridedViewMut<bool>,
1348 a: &StridedView<T, OpA>,
1349 b: &StridedView<T, OpB>,
1350 op: CompareOp,
1351) -> Result<()>
1352where
1353 T: Copy + MaybeSendSync + PartialOrd,
1354 OpA: ElementOp<T>,
1355 OpB: ElementOp<T>,
1356{
1357 match op {
1358 CompareOp::Eq => zip_map2_into(dest, a, b, |lhs, rhs| lhs == rhs),
1359 CompareOp::Lt => zip_map2_into(dest, a, b, |lhs, rhs| lhs < rhs),
1360 CompareOp::Le => zip_map2_into(dest, a, b, |lhs, rhs| lhs <= rhs),
1361 CompareOp::Gt => zip_map2_into(dest, a, b, |lhs, rhs| lhs > rhs),
1362 CompareOp::Ge => zip_map2_into(dest, a, b, |lhs, rhs| lhs >= rhs),
1363 }
1364}
1365
1366pub fn compare_into_uninit<T, OpA, OpB>(
1385 dest: &mut StridedViewMut<MaybeUninit<bool>>,
1386 a: &StridedView<T, OpA>,
1387 b: &StridedView<T, OpB>,
1388 op: CompareOp,
1389) -> Result<()>
1390where
1391 T: Copy + MaybeSendSync + PartialOrd,
1392 OpA: ElementOp<T>,
1393 OpB: ElementOp<T>,
1394{
1395 ensure_same_shape(dest.dims(), a.dims())?;
1396 ensure_same_shape(dest.dims(), b.dims())?;
1397 let validated = validate_destination_layout(dest.dims(), dest.strides())?;
1398 validate_typed_no_overlap(dest, a, 0)?;
1399 validate_typed_no_overlap(dest, b, 1)?;
1400 match op {
1401 CompareOp::Eq => zip_map2_into_validated(
1402 dest,
1403 a,
1404 b,
1405 |lhs, rhs| MaybeUninit::new(lhs == rhs),
1406 validated,
1407 ),
1408 CompareOp::Lt => zip_map2_into_validated(
1409 dest,
1410 a,
1411 b,
1412 |lhs, rhs| MaybeUninit::new(lhs < rhs),
1413 validated,
1414 ),
1415 CompareOp::Le => zip_map2_into_validated(
1416 dest,
1417 a,
1418 b,
1419 |lhs, rhs| MaybeUninit::new(lhs <= rhs),
1420 validated,
1421 ),
1422 CompareOp::Gt => zip_map2_into_validated(
1423 dest,
1424 a,
1425 b,
1426 |lhs, rhs| MaybeUninit::new(lhs > rhs),
1427 validated,
1428 ),
1429 CompareOp::Ge => zip_map2_into_validated(
1430 dest,
1431 a,
1432 b,
1433 |lhs, rhs| MaybeUninit::new(lhs >= rhs),
1434 validated,
1435 ),
1436 }
1437}
1438
1439pub(crate) fn zip_map2_raw_into_validated<
1440 D: Copy + MaybeSendSync,
1441 A: Copy + MaybeSendSync,
1442 B: Copy + MaybeSendSync,
1443 OpA: ElementOp<A>,
1444 OpB: ElementOp<B>,
1445>(
1446 dest: &mut crate::RawStridedMut<'_, D>,
1447 a: &crate::RawStridedRef<'_, A>,
1448 b: &crate::RawStridedRef<'_, B>,
1449 f: impl Fn(A, B) -> D + MaybeSync,
1450 validated: ValidatedDestinationLayout,
1451) -> Result<()> {
1452 ensure_same_shape(dest.dims(), a.dims())?;
1453 ensure_same_shape(dest.dims(), b.dims())?;
1454 zip_map2_parts_into_validated::<D, A, B, OpA, OpB>(
1455 dest.as_mut_ptr(),
1456 dest.dims(),
1457 dest.strides(),
1458 a.ptr(),
1459 a.strides(),
1460 b.ptr(),
1461 b.strides(),
1462 f,
1463 validated,
1464 )
1465}
1466
1467#[allow(clippy::too_many_arguments)]
1468fn zip_map2_parts_into<
1469 D: Copy + MaybeSendSync,
1470 A: Copy + MaybeSendSync,
1471 B: Copy + MaybeSendSync,
1472 OpA: ElementOp<A>,
1473 OpB: ElementOp<B>,
1474>(
1475 dst_ptr: *mut D,
1476 dst_dims: &[usize],
1477 dst_strides: &[isize],
1478 a_ptr: *const A,
1479 a_dims: &[usize],
1480 a_strides: &[isize],
1481 b_ptr: *const B,
1482 b_dims: &[usize],
1483 b_strides: &[isize],
1484 f: impl Fn(A, B) -> D + MaybeSync,
1485) -> Result<()> {
1486 ensure_same_shape(dst_dims, a_dims)?;
1487 ensure_same_shape(dst_dims, b_dims)?;
1488 let validated = validate_destination_layout(dst_dims, dst_strides)?;
1489 zip_map2_parts_into_validated::<D, A, B, OpA, OpB>(
1490 dst_ptr,
1491 dst_dims,
1492 dst_strides,
1493 a_ptr,
1494 a_strides,
1495 b_ptr,
1496 b_strides,
1497 f,
1498 validated,
1499 )
1500}
1501
1502#[allow(clippy::too_many_arguments)]
1503fn zip_map2_parts_into_validated<
1504 D: Copy + MaybeSendSync,
1505 A: Copy + MaybeSendSync,
1506 B: Copy + MaybeSendSync,
1507 OpA: ElementOp<A>,
1508 OpB: ElementOp<B>,
1509>(
1510 dst_ptr: *mut D,
1511 dst_dims: &[usize],
1512 dst_strides: &[isize],
1513 a_ptr: *const A,
1514 a_strides: &[isize],
1515 b_ptr: *const B,
1516 b_strides: &[isize],
1517 f: impl Fn(A, B) -> D + MaybeSync,
1518 _validated: ValidatedDestinationLayout,
1519) -> Result<()> {
1520 if sequential_contiguous_layout(dst_dims, &[dst_strides, a_strides, b_strides]).is_some() {
1521 let len = total_len(dst_dims);
1522 let dst = unsafe { std::slice::from_raw_parts_mut(dst_ptr, len) };
1523 let sa = unsafe { std::slice::from_raw_parts(a_ptr, len) };
1524 let sb = unsafe { std::slice::from_raw_parts(b_ptr, len) };
1525 simd::dispatch_if_large(len, || {
1526 for i in 0..len {
1527 dst[i] = f(OpA::apply(sa[i]), OpB::apply(sb[i]));
1528 }
1529 });
1530 return Ok(());
1531 }
1532
1533 let strides_list: [&[isize]; 3] = [dst_strides, a_strides, b_strides];
1534 let elem_size = std::mem::size_of::<D>()
1535 .max(std::mem::size_of::<A>())
1536 .max(std::mem::size_of::<B>());
1537 let total = total_len(dst_dims);
1538
1539 let (fused_dims, ordered_strides, plan) = if total <= SMALL_TENSOR_THRESHOLD {
1541 build_plan_fused_small(dst_dims, &strides_list)
1542 } else {
1543 build_plan_fused(dst_dims, &strides_list, Some(0), elem_size)
1544 };
1545
1546 #[cfg(feature = "parallel")]
1547 {
1548 let total: usize = fused_dims.iter().product();
1549 let nthreads = crate::execution_policy::rayon_threads();
1550 if total > MINTHREADLENGTH && nthreads > 1 {
1551 use crate::threading::SendPtr;
1552 let dst_send = SendPtr(dst_ptr);
1553 let a_send = SendPtr(a_ptr as *mut A);
1554 let b_send = SendPtr(b_ptr as *mut B);
1555
1556 let costs = compute_costs(&ordered_strides);
1557 let initial_offsets = vec![0isize; strides_list.len()];
1558 return mapreduce_threaded(
1559 &fused_dims,
1560 &plan.block,
1561 &ordered_strides,
1562 &initial_offsets,
1563 &costs,
1564 nthreads,
1565 0,
1566 1,
1567 &|dims, blocks, strides_list, offsets| {
1568 for_each_inner_block_with_offsets(
1569 dims,
1570 blocks,
1571 strides_list,
1572 offsets,
1573 |offsets, len, strides| {
1574 let dp = unsafe { dst_send.as_ptr().offset(offsets[0]) };
1575 let ap = unsafe { a_send.as_const().offset(offsets[1]) };
1576 let bp = unsafe { b_send.as_const().offset(offsets[2]) };
1577 unsafe {
1578 inner_loop_map2::<D, A, B, OpA, OpB>(
1579 dp, strides[0], ap, strides[1], bp, strides[2], len, &f,
1580 )
1581 };
1582 Ok(())
1583 },
1584 )
1585 },
1586 );
1587 }
1588 }
1589
1590 let initial_offsets = vec![0isize; ordered_strides.len()];
1591 for_each_inner_block_preordered(
1592 &fused_dims,
1593 &plan.block,
1594 &ordered_strides,
1595 &initial_offsets,
1596 |offsets, len, strides| {
1597 let dp = unsafe { dst_ptr.offset(offsets[0]) };
1598 let ap = unsafe { a_ptr.offset(offsets[1]) };
1599 let bp = unsafe { b_ptr.offset(offsets[2]) };
1600 unsafe {
1601 inner_loop_map2::<D, A, B, OpA, OpB>(
1602 dp, strides[0], ap, strides[1], bp, strides[2], len, &f,
1603 )
1604 };
1605 Ok(())
1606 },
1607 )
1608}
1609
1610fn mul_identity_into_raw<
1611 O: MulOutput<D>,
1612 D: Copy + MaybeSendSync + 'static,
1613 A: Copy + MaybeSendSync + Mul<B, Output = D> + 'static,
1614 B: Copy + MaybeSendSync + 'static,
1615>(
1616 dst_ptr: *mut O::Slot,
1617 dst_dims: &[usize],
1618 dst_strides: &[isize],
1619 a_ptr: *const A,
1620 a_strides: &[isize],
1621 b_ptr: *const B,
1622 b_strides: &[isize],
1623 _validated: ValidatedDestinationLayout,
1624) -> Result<()> {
1625 debug_assert_eq!(dst_dims.len(), a_strides.len());
1626 debug_assert_eq!(dst_dims.len(), b_strides.len());
1627
1628 if sequential_contiguous_layout(dst_dims, &[dst_strides, a_strides, b_strides]).is_some() {
1629 let len = total_len(dst_dims);
1630 let sa = unsafe { std::slice::from_raw_parts(a_ptr, len) };
1631 let sb = unsafe { std::slice::from_raw_parts(b_ptr, len) };
1632 if unsafe { O::try_contiguous(dst_ptr, len, sa, sb) } {
1633 return Ok(());
1634 }
1635 for i in 0..len {
1636 unsafe { O::write(dst_ptr.add(i), multiply_value(sa[i], sb[i])) };
1637 }
1638 return Ok(());
1639 }
1640
1641 let strides_list: [&[isize]; 3] = [dst_strides, a_strides, b_strides];
1642 let elem_size = std::mem::size_of::<D>()
1643 .max(std::mem::size_of::<A>())
1644 .max(std::mem::size_of::<B>());
1645 let total = total_len(dst_dims);
1646
1647 if try_contiguous_range_mul::<O, D, A, B>(
1648 dst_ptr,
1649 dst_dims,
1650 dst_strides,
1651 a_ptr,
1652 a_strides,
1653 b_ptr,
1654 b_strides,
1655 ) {
1656 return Ok(());
1657 }
1658
1659 let (fused_dims, ordered_strides, plan) = if total <= SMALL_TENSOR_THRESHOLD {
1660 build_plan_fused_small(dst_dims, &strides_list)
1661 } else {
1662 build_plan_fused(dst_dims, &strides_list, Some(0), elem_size)
1663 };
1664
1665 #[cfg(feature = "parallel")]
1666 {
1667 let total: usize = fused_dims.iter().product();
1668 let nthreads = crate::execution_policy::rayon_threads();
1669 if total > MINTHREADLENGTH && nthreads > 1 {
1670 use crate::threading::SendPtr;
1671 let dst_send = SendPtr(dst_ptr);
1672 let a_send = SendPtr(a_ptr as *mut A);
1673 let b_send = SendPtr(b_ptr as *mut B);
1674
1675 let costs = compute_costs(&ordered_strides);
1676 let initial_offsets = vec![0isize; strides_list.len()];
1677 return mapreduce_threaded(
1678 &fused_dims,
1679 &plan.block,
1680 &ordered_strides,
1681 &initial_offsets,
1682 &costs,
1683 nthreads,
1684 0,
1685 1,
1686 &|dims, blocks, strides_list, offsets| {
1687 for_each_inner_block_with_offsets(
1688 dims,
1689 blocks,
1690 strides_list,
1691 offsets,
1692 |offsets, len, strides| {
1693 let dp = unsafe { dst_send.as_ptr().offset(offsets[0]) };
1694 let ap = unsafe { a_send.as_const().offset(offsets[1]) };
1695 let bp = unsafe { b_send.as_const().offset(offsets[2]) };
1696 unsafe {
1697 inner_loop_mul2::<O, D, A, B>(
1698 dp, strides[0], ap, strides[1], bp, strides[2], len,
1699 )
1700 };
1701 Ok(())
1702 },
1703 )
1704 },
1705 );
1706 }
1707 }
1708
1709 let initial_offsets = vec![0isize; ordered_strides.len()];
1710 for_each_inner_block_preordered(
1711 &fused_dims,
1712 &plan.block,
1713 &ordered_strides,
1714 &initial_offsets,
1715 |offsets, len, strides| {
1716 let dp = unsafe { dst_ptr.offset(offsets[0]) };
1717 let ap = unsafe { a_ptr.offset(offsets[1]) };
1718 let bp = unsafe { b_ptr.offset(offsets[2]) };
1719 unsafe {
1720 inner_loop_mul2::<O, D, A, B>(dp, strides[0], ap, strides[1], bp, strides[2], len)
1721 };
1722 Ok(())
1723 },
1724 )
1725}
1726
1727pub fn mul_into<
1732 D: Copy + MaybeSendSync + 'static,
1733 A: Copy + Mul<B, Output = D> + MaybeSendSync + 'static,
1734 B: Copy + MaybeSendSync + 'static,
1735 OpA: ElementOp<A>,
1736 OpB: ElementOp<B>,
1737>(
1738 dest: &mut StridedViewMut<D>,
1739 a: &StridedView<A, OpA>,
1740 b: &StridedView<B, OpB>,
1741) -> Result<()> {
1742 ensure_same_shape(dest.dims(), a.dims())?;
1743 ensure_same_shape(dest.dims(), b.dims())?;
1744 let validated = validate_destination_layout(dest.dims(), dest.strides())?;
1745
1746 if OpA::IS_IDENTITY && OpB::IS_IDENTITY {
1747 return mul_identity_into_raw::<InitializedOutput, D, A, B>(
1748 dest.as_mut_ptr(),
1749 dest.dims(),
1750 dest.strides(),
1751 a.ptr(),
1752 a.strides(),
1753 b.ptr(),
1754 b.strides(),
1755 validated,
1756 );
1757 }
1758
1759 zip_map2_into_validated(dest, a, b, multiply_value, validated)
1760}
1761
1762pub fn mul_into_uninit<
1781 D: Copy + MaybeSendSync + 'static,
1782 A: Copy + Mul<B, Output = D> + MaybeSendSync + 'static,
1783 B: Copy + MaybeSendSync + 'static,
1784 OpA: ElementOp<A>,
1785 OpB: ElementOp<B>,
1786>(
1787 dest: &mut StridedViewMut<MaybeUninit<D>>,
1788 a: &StridedView<A, OpA>,
1789 b: &StridedView<B, OpB>,
1790) -> Result<()> {
1791 ensure_same_shape(dest.dims(), a.dims())?;
1792 ensure_same_shape(dest.dims(), b.dims())?;
1793 let validated = validate_destination_layout(dest.dims(), dest.strides())?;
1794 validate_typed_no_overlap(dest, a, 0)?;
1795 validate_typed_no_overlap(dest, b, 1)?;
1796
1797 if OpA::IS_IDENTITY && OpB::IS_IDENTITY {
1798 return mul_identity_into_raw::<UninitializedOutput, D, A, B>(
1799 dest.as_mut_ptr(),
1800 dest.dims(),
1801 dest.strides(),
1802 a.ptr(),
1803 a.strides(),
1804 b.ptr(),
1805 b.strides(),
1806 validated,
1807 );
1808 }
1809
1810 zip_map2_into_validated(
1811 dest,
1812 a,
1813 b,
1814 |lhs, rhs| MaybeUninit::new(multiply_value(lhs, rhs)),
1815 validated,
1816 )
1817}
1818
1819fn broadcast_strides_for_axes(
1820 source_dims: &[usize],
1821 source_strides: &[isize],
1822 target_dims: &[usize],
1823 axes: &[usize],
1824) -> Result<AxisVec<isize>> {
1825 if source_dims.len() != axes.len() {
1826 return Err(StridedError::RankMismatch(source_dims.len(), axes.len()));
1827 }
1828 debug_assert_eq!(source_dims.len(), source_strides.len());
1829
1830 let mut seen = AxisVec::<bool>::new();
1831 seen.resize(target_dims.len(), false);
1832 let mut strides = AxisVec::<isize>::new();
1833 strides.resize(target_dims.len(), 0);
1834 for (src_axis, &dst_axis) in axes.iter().enumerate() {
1835 if dst_axis >= target_dims.len() {
1836 return Err(StridedError::InvalidAxis {
1837 axis: dst_axis,
1838 rank: target_dims.len(),
1839 });
1840 }
1841 if seen[dst_axis] {
1842 return Err(StridedError::InvalidAxis {
1843 axis: dst_axis,
1844 rank: target_dims.len(),
1845 });
1846 }
1847 seen[dst_axis] = true;
1848
1849 let source_dim = source_dims[src_axis];
1850 let target_dim = target_dims[dst_axis];
1851 if source_dim != target_dim && source_dim != 1 {
1852 return Err(StridedError::ShapeMismatch(
1853 source_dims.to_vec(),
1854 target_dims.to_vec(),
1855 ));
1856 }
1857 if source_dim == target_dim {
1858 strides[dst_axis] = source_strides[src_axis];
1859 }
1860 }
1861
1862 Ok(strides)
1863}
1864
1865fn broadcast_view_with_strides<'a, T, Op: ElementOp<T>>(
1866 view: &StridedView<'a, T, Op>,
1867 target_dims: &[usize],
1868 strides: &[isize],
1869) -> StridedView<'a, T, Op> {
1870 unsafe { StridedView::new_unchecked(view.data(), target_dims, strides, view.offset()) }
1871}
1872
1873pub fn broadcast_mul_into<
1878 D: Copy + MaybeSendSync + 'static,
1879 A: Copy + Mul<B, Output = D> + MaybeSendSync + 'static,
1880 B: Copy + MaybeSendSync + 'static,
1881 OpA: ElementOp<A>,
1882 OpB: ElementOp<B>,
1883>(
1884 dest: &mut StridedViewMut<D>,
1885 a: &StridedView<A, OpA>,
1886 a_axes: &[usize],
1887 b: &StridedView<B, OpB>,
1888 b_axes: &[usize],
1889) -> Result<()> {
1890 let validated = validate_destination_layout(dest.dims(), dest.strides())?;
1891 let a_strides = broadcast_strides_for_axes(a.dims(), a.strides(), dest.dims(), a_axes)?;
1892 let b_strides = broadcast_strides_for_axes(b.dims(), b.strides(), dest.dims(), b_axes)?;
1893
1894 if OpA::IS_IDENTITY && OpB::IS_IDENTITY {
1895 return mul_identity_into_raw::<InitializedOutput, D, A, B>(
1896 dest.as_mut_ptr(),
1897 dest.dims(),
1898 dest.strides(),
1899 a.ptr(),
1900 &a_strides,
1901 b.ptr(),
1902 &b_strides,
1903 validated,
1904 );
1905 }
1906
1907 let a = broadcast_view_with_strides(a, dest.dims(), &a_strides);
1908 let b = broadcast_view_with_strides(b, dest.dims(), &b_strides);
1909 zip_map2_into_validated(dest, &a, &b, multiply_value, validated)
1910}
1911
1912pub fn broadcast_mul_into_uninit<
1931 D: Copy + MaybeSendSync + 'static,
1932 A: Copy + Mul<B, Output = D> + MaybeSendSync + 'static,
1933 B: Copy + MaybeSendSync + 'static,
1934 OpA: ElementOp<A>,
1935 OpB: ElementOp<B>,
1936>(
1937 dest: &mut StridedViewMut<MaybeUninit<D>>,
1938 a: &StridedView<A, OpA>,
1939 a_axes: &[usize],
1940 b: &StridedView<B, OpB>,
1941 b_axes: &[usize],
1942) -> Result<()> {
1943 let validated = validate_destination_layout(dest.dims(), dest.strides())?;
1944 let a_strides = broadcast_strides_for_axes(a.dims(), a.strides(), dest.dims(), a_axes)?;
1945 let b_strides = broadcast_strides_for_axes(b.dims(), b.strides(), dest.dims(), b_axes)?;
1946 validate_typed_no_overlap(dest, a, 0)?;
1947 validate_typed_no_overlap(dest, b, 1)?;
1948
1949 if OpA::IS_IDENTITY && OpB::IS_IDENTITY {
1950 return mul_identity_into_raw::<UninitializedOutput, D, A, B>(
1951 dest.as_mut_ptr(),
1952 dest.dims(),
1953 dest.strides(),
1954 a.ptr(),
1955 &a_strides,
1956 b.ptr(),
1957 &b_strides,
1958 validated,
1959 );
1960 }
1961
1962 let a = broadcast_view_with_strides(a, dest.dims(), &a_strides);
1963 let b = broadcast_view_with_strides(b, dest.dims(), &b_strides);
1964 zip_map2_into_validated(
1965 dest,
1966 &a,
1967 &b,
1968 |lhs, rhs| MaybeUninit::new(multiply_value(lhs, rhs)),
1969 validated,
1970 )
1971}
1972
1973pub fn zip_map3_into<
1975 D: Copy + MaybeSendSync,
1976 A: Copy + MaybeSendSync,
1977 B: Copy + MaybeSendSync,
1978 C: Copy + MaybeSendSync,
1979 OpA: ElementOp<A>,
1980 OpB: ElementOp<B>,
1981 OpC: ElementOp<C>,
1982>(
1983 dest: &mut StridedViewMut<D>,
1984 a: &StridedView<A, OpA>,
1985 b: &StridedView<B, OpB>,
1986 c: &StridedView<C, OpC>,
1987 f: impl Fn(A, B, C) -> D + MaybeSync,
1988) -> Result<()> {
1989 ensure_same_shape(dest.dims(), a.dims())?;
1990 ensure_same_shape(dest.dims(), b.dims())?;
1991 ensure_same_shape(dest.dims(), c.dims())?;
1992 let validated = validate_destination_layout(dest.dims(), dest.strides())?;
1993 zip_map3_into_validated(dest, a, b, c, f, validated)
1994}
1995
1996pub(crate) fn zip_map3_into_validated<
1997 D: Copy + MaybeSendSync,
1998 A: Copy + MaybeSendSync,
1999 B: Copy + MaybeSendSync,
2000 C: Copy + MaybeSendSync,
2001 OpA: ElementOp<A>,
2002 OpB: ElementOp<B>,
2003 OpC: ElementOp<C>,
2004>(
2005 dest: &mut StridedViewMut<D>,
2006 a: &StridedView<A, OpA>,
2007 b: &StridedView<B, OpB>,
2008 c: &StridedView<C, OpC>,
2009 f: impl Fn(A, B, C) -> D + MaybeSync,
2010 _validated: ValidatedDestinationLayout,
2011) -> Result<()> {
2012 ensure_same_shape(dest.dims(), a.dims())?;
2013 ensure_same_shape(dest.dims(), b.dims())?;
2014 ensure_same_shape(dest.dims(), c.dims())?;
2015 let dst_ptr = dest.as_mut_ptr();
2016 let a_ptr = a.ptr();
2017 let b_ptr = b.ptr();
2018 let c_ptr = c.ptr();
2019
2020 let dst_dims = dest.dims();
2021 let dst_strides = dest.strides();
2022
2023 if sequential_contiguous_layout(
2024 dst_dims,
2025 &[dst_strides, a.strides(), b.strides(), c.strides()],
2026 )
2027 .is_some()
2028 {
2029 let len = total_len(dst_dims);
2030 let dst = unsafe { std::slice::from_raw_parts_mut(dst_ptr, len) };
2031 let sa = unsafe { std::slice::from_raw_parts(a_ptr, len) };
2032 let sb = unsafe { std::slice::from_raw_parts(b_ptr, len) };
2033 let sc = unsafe { std::slice::from_raw_parts(c_ptr, len) };
2034 simd::dispatch_if_large(len, || {
2035 for i in 0..len {
2036 dst[i] = f(OpA::apply(sa[i]), OpB::apply(sb[i]), OpC::apply(sc[i]));
2037 }
2038 });
2039 return Ok(());
2040 }
2041
2042 let strides_list: [&[isize]; 4] = [dst_strides, a.strides(), b.strides(), c.strides()];
2043 let elem_size = std::mem::size_of::<D>()
2044 .max(std::mem::size_of::<A>())
2045 .max(std::mem::size_of::<B>())
2046 .max(std::mem::size_of::<C>());
2047 let total = total_len(dst_dims);
2048
2049 let (fused_dims, ordered_strides, plan) = if total <= SMALL_TENSOR_THRESHOLD {
2051 build_plan_fused_small(dst_dims, &strides_list)
2052 } else {
2053 build_plan_fused(dst_dims, &strides_list, Some(0), elem_size)
2054 };
2055
2056 #[cfg(feature = "parallel")]
2057 {
2058 let total: usize = fused_dims.iter().product();
2059 let nthreads = crate::execution_policy::rayon_threads();
2060 if total > MINTHREADLENGTH && nthreads > 1 {
2061 use crate::threading::SendPtr;
2062 let dst_send = SendPtr(dst_ptr);
2063 let a_send = SendPtr(a_ptr as *mut A);
2064 let b_send = SendPtr(b_ptr as *mut B);
2065 let c_send = SendPtr(c_ptr as *mut C);
2066
2067 let costs = compute_costs(&ordered_strides);
2068 let initial_offsets = vec![0isize; strides_list.len()];
2069 return mapreduce_threaded(
2070 &fused_dims,
2071 &plan.block,
2072 &ordered_strides,
2073 &initial_offsets,
2074 &costs,
2075 nthreads,
2076 0,
2077 1,
2078 &|dims, blocks, strides_list, offsets| {
2079 for_each_inner_block_with_offsets(
2080 dims,
2081 blocks,
2082 strides_list,
2083 offsets,
2084 |offsets, len, strides| {
2085 let dp = unsafe { dst_send.as_ptr().offset(offsets[0]) };
2086 let ap = unsafe { a_send.as_const().offset(offsets[1]) };
2087 let bp = unsafe { b_send.as_const().offset(offsets[2]) };
2088 let cp = unsafe { c_send.as_const().offset(offsets[3]) };
2089 unsafe {
2090 inner_loop_map3::<D, A, B, C, OpA, OpB, OpC>(
2091 dp, strides[0], ap, strides[1], bp, strides[2], cp, strides[3],
2092 len, &f,
2093 )
2094 };
2095 Ok(())
2096 },
2097 )
2098 },
2099 );
2100 }
2101 }
2102
2103 let initial_offsets = vec![0isize; ordered_strides.len()];
2104 for_each_inner_block_preordered(
2105 &fused_dims,
2106 &plan.block,
2107 &ordered_strides,
2108 &initial_offsets,
2109 |offsets, len, strides| {
2110 let dp = unsafe { dst_ptr.offset(offsets[0]) };
2111 let ap = unsafe { a_ptr.offset(offsets[1]) };
2112 let bp = unsafe { b_ptr.offset(offsets[2]) };
2113 let cp = unsafe { c_ptr.offset(offsets[3]) };
2114 unsafe {
2115 inner_loop_map3::<D, A, B, C, OpA, OpB, OpC>(
2116 dp, strides[0], ap, strides[1], bp, strides[2], cp, strides[3], len, &f,
2117 )
2118 };
2119 Ok(())
2120 },
2121 )
2122}
2123
2124pub fn zip_map4_into<
2126 D: Copy + MaybeSendSync,
2127 A: Copy + MaybeSendSync,
2128 B: Copy + MaybeSendSync,
2129 C: Copy + MaybeSendSync,
2130 E: Copy + MaybeSendSync,
2131 OpA: ElementOp<A>,
2132 OpB: ElementOp<B>,
2133 OpC: ElementOp<C>,
2134 OpE: ElementOp<E>,
2135>(
2136 dest: &mut StridedViewMut<D>,
2137 a: &StridedView<A, OpA>,
2138 b: &StridedView<B, OpB>,
2139 c: &StridedView<C, OpC>,
2140 e: &StridedView<E, OpE>,
2141 f: impl Fn(A, B, C, E) -> D + MaybeSync,
2142) -> Result<()> {
2143 ensure_same_shape(dest.dims(), a.dims())?;
2144 ensure_same_shape(dest.dims(), b.dims())?;
2145 ensure_same_shape(dest.dims(), c.dims())?;
2146 ensure_same_shape(dest.dims(), e.dims())?;
2147 let validated = validate_destination_layout(dest.dims(), dest.strides())?;
2148 zip_map4_into_validated(dest, a, b, c, e, f, validated)
2149}
2150
2151pub(crate) fn zip_map4_into_validated<
2152 D: Copy + MaybeSendSync,
2153 A: Copy + MaybeSendSync,
2154 B: Copy + MaybeSendSync,
2155 C: Copy + MaybeSendSync,
2156 E: Copy + MaybeSendSync,
2157 OpA: ElementOp<A>,
2158 OpB: ElementOp<B>,
2159 OpC: ElementOp<C>,
2160 OpE: ElementOp<E>,
2161>(
2162 dest: &mut StridedViewMut<D>,
2163 a: &StridedView<A, OpA>,
2164 b: &StridedView<B, OpB>,
2165 c: &StridedView<C, OpC>,
2166 e: &StridedView<E, OpE>,
2167 f: impl Fn(A, B, C, E) -> D + MaybeSync,
2168 _validated: ValidatedDestinationLayout,
2169) -> Result<()> {
2170 ensure_same_shape(dest.dims(), a.dims())?;
2171 ensure_same_shape(dest.dims(), b.dims())?;
2172 ensure_same_shape(dest.dims(), c.dims())?;
2173 ensure_same_shape(dest.dims(), e.dims())?;
2174 let dst_ptr = dest.as_mut_ptr();
2175 let a_ptr = a.ptr();
2176 let b_ptr = b.ptr();
2177 let c_ptr = c.ptr();
2178 let e_ptr = e.ptr();
2179
2180 let dst_dims = dest.dims();
2181 let dst_strides = dest.strides();
2182
2183 if sequential_contiguous_layout(
2184 dst_dims,
2185 &[
2186 dst_strides,
2187 a.strides(),
2188 b.strides(),
2189 c.strides(),
2190 e.strides(),
2191 ],
2192 )
2193 .is_some()
2194 {
2195 let len = total_len(dst_dims);
2196 let dst = unsafe { std::slice::from_raw_parts_mut(dst_ptr, len) };
2197 let sa = unsafe { std::slice::from_raw_parts(a_ptr, len) };
2198 let sb = unsafe { std::slice::from_raw_parts(b_ptr, len) };
2199 let sc = unsafe { std::slice::from_raw_parts(c_ptr, len) };
2200 let se = unsafe { std::slice::from_raw_parts(e_ptr, len) };
2201 simd::dispatch_if_large(len, || {
2202 for i in 0..len {
2203 dst[i] = f(
2204 OpA::apply(sa[i]),
2205 OpB::apply(sb[i]),
2206 OpC::apply(sc[i]),
2207 OpE::apply(se[i]),
2208 );
2209 }
2210 });
2211 return Ok(());
2212 }
2213
2214 let strides_list: [&[isize]; 5] = [
2215 dst_strides,
2216 a.strides(),
2217 b.strides(),
2218 c.strides(),
2219 e.strides(),
2220 ];
2221 let elem_size = std::mem::size_of::<D>()
2222 .max(std::mem::size_of::<A>())
2223 .max(std::mem::size_of::<B>())
2224 .max(std::mem::size_of::<C>())
2225 .max(std::mem::size_of::<E>());
2226 let total = total_len(dst_dims);
2227
2228 let (fused_dims, ordered_strides, plan) = if total <= SMALL_TENSOR_THRESHOLD {
2230 build_plan_fused_small(dst_dims, &strides_list)
2231 } else {
2232 build_plan_fused(dst_dims, &strides_list, Some(0), elem_size)
2233 };
2234
2235 #[cfg(feature = "parallel")]
2236 {
2237 let total: usize = fused_dims.iter().product();
2238 let nthreads = crate::execution_policy::rayon_threads();
2239 if total > MINTHREADLENGTH && nthreads > 1 {
2240 use crate::threading::SendPtr;
2241 let dst_send = SendPtr(dst_ptr);
2242 let a_send = SendPtr(a_ptr as *mut A);
2243 let b_send = SendPtr(b_ptr as *mut B);
2244 let c_send = SendPtr(c_ptr as *mut C);
2245 let e_send = SendPtr(e_ptr as *mut E);
2246
2247 let costs = compute_costs(&ordered_strides);
2248 let initial_offsets = vec![0isize; strides_list.len()];
2249 return mapreduce_threaded(
2250 &fused_dims,
2251 &plan.block,
2252 &ordered_strides,
2253 &initial_offsets,
2254 &costs,
2255 nthreads,
2256 0,
2257 1,
2258 &|dims, blocks, strides_list, offsets| {
2259 for_each_inner_block_with_offsets(
2260 dims,
2261 blocks,
2262 strides_list,
2263 offsets,
2264 |offsets, len, strides| {
2265 let dp = unsafe { dst_send.as_ptr().offset(offsets[0]) };
2266 let ap = unsafe { a_send.as_const().offset(offsets[1]) };
2267 let bp = unsafe { b_send.as_const().offset(offsets[2]) };
2268 let cp = unsafe { c_send.as_const().offset(offsets[3]) };
2269 let ep = unsafe { e_send.as_const().offset(offsets[4]) };
2270 unsafe {
2271 inner_loop_map4::<D, A, B, C, E, OpA, OpB, OpC, OpE>(
2272 dp, strides[0], ap, strides[1], bp, strides[2], cp, strides[3],
2273 ep, strides[4], len, &f,
2274 )
2275 };
2276 Ok(())
2277 },
2278 )
2279 },
2280 );
2281 }
2282 }
2283
2284 let initial_offsets = vec![0isize; ordered_strides.len()];
2285 for_each_inner_block_preordered(
2286 &fused_dims,
2287 &plan.block,
2288 &ordered_strides,
2289 &initial_offsets,
2290 |offsets, len, strides| {
2291 let dp = unsafe { dst_ptr.offset(offsets[0]) };
2292 let ap = unsafe { a_ptr.offset(offsets[1]) };
2293 let bp = unsafe { b_ptr.offset(offsets[2]) };
2294 let cp = unsafe { c_ptr.offset(offsets[3]) };
2295 let ep = unsafe { e_ptr.offset(offsets[4]) };
2296 unsafe {
2297 inner_loop_map4::<D, A, B, C, E, OpA, OpB, OpC, OpE>(
2298 dp, strides[0], ap, strides[1], bp, strides[2], cp, strides[3], ep, strides[4],
2299 len, &f,
2300 )
2301 };
2302 Ok(())
2303 },
2304 )
2305}
2306
2307#[cfg(test)]
2308#[path = "map_view/tests/scalar_branch_tests.rs"]
2309mod scalar_branch_tests;
2310
2311#[cfg(all(test, feature = "parallel"))]
2312#[path = "map_view/tests/tests.rs"]
2313mod tests;