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
454impl ContiguousMulRangePlan {
455 fn walks_inputs_without_blocking(&self) -> bool {
465 let streams = |stride: isize| stride == 0 || stride == 1;
466 if streams(self.a_fast_stride) && streams(self.b_fast_stride) {
467 return true;
468 }
469 #[cfg(feature = "parallel")]
470 if transposed_scalar_tile_kind(self).is_some() {
471 return true;
472 }
473 false
474 }
475}
476
477#[cfg(feature = "parallel")]
478#[derive(Clone, Copy, Debug, Eq, PartialEq)]
479enum TransposedScalarTileKind {
480 RhsScalar,
481 LhsScalar,
482}
483
484#[cfg(feature = "parallel")]
485fn transposed_scalar_tile_kind(plan: &ContiguousMulRangePlan) -> Option<TransposedScalarTileKind> {
486 let row_len = isize::try_from(plan.row_len).ok()?;
487 if plan.a_fast_stride == row_len
488 && plan.a_row_stride == 1
489 && plan.b_fast_stride == 0
490 && plan.b_row_stride == 0
491 {
492 return Some(TransposedScalarTileKind::RhsScalar);
493 }
494
495 if plan.b_fast_stride == row_len
496 && plan.b_row_stride == 1
497 && plan.a_fast_stride == 0
498 && plan.a_row_stride == 0
499 {
500 return Some(TransposedScalarTileKind::LhsScalar);
501 }
502
503 None
504}
505
506fn compact_axis_order(dims: &[usize], strides: &[isize]) -> Option<AxisVec<usize>> {
507 if dims.len() != strides.len() {
508 return None;
509 }
510
511 let mut active = AxisVec::<usize>::new();
512 let mut inactive = AxisVec::<usize>::new();
513 for (axis, (&dim, &stride)) in dims.iter().zip(strides.iter()).enumerate() {
514 if stride < 0 {
515 return None;
516 }
517 if dim > 1 {
518 active.push(axis);
519 } else {
520 inactive.push(axis);
521 }
522 }
523
524 active.sort_by(|&lhs, &rhs| strides[lhs].cmp(&strides[rhs]).then_with(|| lhs.cmp(&rhs)));
525
526 let mut expected = 1isize;
527 for &axis in &active {
528 if strides[axis] != expected {
529 return None;
530 }
531 expected = expected.saturating_mul(dims[axis] as isize);
532 }
533
534 active.extend(inactive);
535 Some(active)
536}
537
538fn can_fuse_contiguous_range_axis(dim: usize, prev_stride: isize, next_stride: isize) -> bool {
539 dim <= 1 || (prev_stride == 0 && next_stride == 0) || next_stride == prev_stride * dim as isize
540}
541
542fn contiguous_mul_range_plan(
543 dims: &[usize],
544 dst_strides: &[isize],
545 a_strides: &[isize],
546 b_strides: &[isize],
547) -> Option<ContiguousMulRangePlan> {
548 let axis_order = compact_axis_order(dims, dst_strides)?;
549 if dims.is_empty() {
550 return Some(ContiguousMulRangePlan {
551 axis_order,
552 inner_len: 1,
553 inner_axis_count: 0,
554 row_len: 1,
555 outer_axis_start: 0,
556 fast_axis: 0,
557 a_fast_stride: 0,
558 b_fast_stride: 0,
559 a_row_stride: 0,
560 b_row_stride: 0,
561 });
562 }
563
564 let first_pos = axis_order
565 .iter()
566 .position(|&axis| dims[axis] > 1)
567 .unwrap_or(0);
568 let first_axis = axis_order[first_pos];
569 let mut inner_len = dims[first_axis].max(1);
570 let mut inner_axis_count = first_pos + 1;
571 let mut prev_axis = first_axis;
572
573 for &axis in axis_order.iter().skip(first_pos + 1) {
574 if can_fuse_contiguous_range_axis(
575 dims[prev_axis],
576 dst_strides[prev_axis],
577 dst_strides[axis],
578 ) && can_fuse_contiguous_range_axis(
579 dims[prev_axis],
580 a_strides[prev_axis],
581 a_strides[axis],
582 ) && can_fuse_contiguous_range_axis(
583 dims[prev_axis],
584 b_strides[prev_axis],
585 b_strides[axis],
586 ) {
587 inner_len = inner_len.checked_mul(dims[axis].max(1))?;
588 inner_axis_count += 1;
589 prev_axis = axis;
590 } else {
591 break;
592 }
593 }
594
595 let row = axis_order
596 .iter()
597 .enumerate()
598 .skip(inner_axis_count)
599 .find(|&(_, &axis)| dims[axis] > 1);
600 let (row_len, outer_axis_start, a_row_stride, b_row_stride) =
601 if let Some((row_pos, &row_axis)) = row {
602 (
603 dims[row_axis],
604 row_pos + 1,
605 a_strides[row_axis],
606 b_strides[row_axis],
607 )
608 } else {
609 (1, axis_order.len(), 0, 0)
610 };
611
612 Some(ContiguousMulRangePlan {
613 axis_order,
614 inner_len,
615 inner_axis_count,
616 row_len,
617 outer_axis_start,
618 fast_axis: first_axis,
619 a_fast_stride: a_strides[first_axis],
620 b_fast_stride: b_strides[first_axis],
621 a_row_stride,
622 b_row_stride,
623 })
624}
625
626struct ContiguousMulOuterCursor<'a> {
627 dims: &'a [usize],
628 a_strides: &'a [isize],
629 b_strides: &'a [isize],
630 axes: AxisVec<usize>,
631 coords: AxisVec<usize>,
632 a_offset: isize,
633 b_offset: isize,
634}
635
636impl<'a> ContiguousMulOuterCursor<'a> {
637 fn new(
638 dims: &'a [usize],
639 a_strides: &'a [isize],
640 b_strides: &'a [isize],
641 plan: &ContiguousMulRangePlan,
642 outer_group: usize,
643 ) -> Self {
644 let axes: AxisVec<usize> = plan
645 .axis_order
646 .iter()
647 .skip(plan.outer_axis_start)
648 .copied()
649 .collect();
650 let mut coords = AxisVec::<usize>::with_capacity(axes.len());
651 let mut rem = outer_group;
652 let mut a_offset = 0isize;
653 let mut b_offset = 0isize;
654
655 for &axis in &axes {
656 let dim = dims[axis].max(1);
657 let coord = rem % dim;
658 rem /= dim;
659 coords.push(coord);
660 a_offset += coord as isize * a_strides[axis];
661 b_offset += coord as isize * b_strides[axis];
662 }
663
664 Self {
665 dims,
666 a_strides,
667 b_strides,
668 axes,
669 coords,
670 a_offset,
671 b_offset,
672 }
673 }
674
675 fn advance(&mut self) {
676 for (i, &axis) in self.axes.iter().enumerate() {
677 let dim = self.dims[axis].max(1);
678 if dim <= 1 {
679 continue;
680 }
681
682 if self.coords[i] + 1 < dim {
685 self.coords[i] += 1;
686 self.a_offset += self.a_strides[axis];
687 self.b_offset += self.b_strides[axis];
688 break;
689 }
690
691 let last = (dim - 1) as isize;
692 self.coords[i] = 0;
693 self.a_offset -= last * self.a_strides[axis];
694 self.b_offset -= last * self.b_strides[axis];
695 }
696 }
697}
698
699#[inline(always)]
700unsafe fn run_contiguous_mul_row_block<
701 O: MulOutput<D>,
702 D: Copy + 'static,
703 A: Copy + Mul<B, Output = D> + 'static,
704 B: Copy + 'static,
705>(
706 dst_ptr: *mut O::Slot,
707 a_ptr: *const A,
708 b_ptr: *const B,
709 plan: &ContiguousMulRangePlan,
710 base_index: usize,
711 total: usize,
712 base_a_offset: isize,
713 base_b_offset: isize,
714) {
715 let inner_len = plan.inner_len.max(1);
716 let row_len = plan.row_len.max(1);
717 #[cfg(feature = "parallel")]
718 let block_len = inner_len.saturating_mul(row_len);
719
720 #[cfg(feature = "parallel")]
721 if total.saturating_sub(base_index) >= block_len {
722 match transposed_scalar_tile_kind(plan) {
723 Some(TransposedScalarTileKind::RhsScalar) => {
724 if simd::try_mul_transposed_scalar_rhs_2d::<D, A, B>(
725 dst_ptr.add(base_index).cast::<D>(),
726 a_ptr.offset(base_a_offset),
727 b_ptr.offset(base_b_offset),
728 inner_len,
729 row_len,
730 plan.a_fast_stride,
731 plan.a_row_stride,
732 ) {
733 return;
734 }
735 }
736 Some(TransposedScalarTileKind::LhsScalar) => {
737 if simd::try_mul_transposed_scalar_lhs_2d::<D, A, B>(
738 dst_ptr.add(base_index).cast::<D>(),
739 a_ptr.offset(base_a_offset),
740 b_ptr.offset(base_b_offset),
741 inner_len,
742 row_len,
743 plan.b_fast_stride,
744 plan.b_row_stride,
745 ) {
746 return;
747 }
748 }
749 None => {}
750 }
751 }
752
753 let mut index = base_index;
754 let mut a_offset = base_a_offset;
755 let mut b_offset = base_b_offset;
756
757 for _ in 0..row_len {
758 if index >= total {
759 break;
760 }
761 let len = inner_len.min(total - index);
762 inner_loop_mul2::<O, D, A, B>(
763 dst_ptr.add(index),
764 1,
765 a_ptr.offset(a_offset),
766 plan.a_fast_stride,
767 b_ptr.offset(b_offset),
768 plan.b_fast_stride,
769 len,
770 );
771 index += inner_len;
772 a_offset += plan.a_row_stride;
773 b_offset += plan.b_row_stride;
774 }
775}
776
777#[cfg(feature = "parallel")]
778fn strided_offset_for_contiguous_linear_index(
779 dims: &[usize],
780 strides: &[isize],
781 axis_order: &[usize],
782 mut index: usize,
783) -> isize {
784 let mut offset = 0isize;
785 for &axis in axis_order {
786 let dim = dims[axis];
787 if dim == 0 {
788 return 0;
789 }
790 let coord = index % dim;
791 index /= dim;
792 offset += coord as isize * strides[axis];
793 }
794 offset
795}
796
797fn try_contiguous_range_mul<
798 O: MulOutput<D>,
799 D: Copy + MaybeSendSync + 'static,
800 A: Copy + MaybeSendSync + Mul<B, Output = D> + 'static,
801 B: Copy + MaybeSendSync + 'static,
802>(
803 dst_ptr: *mut O::Slot,
804 dims: &[usize],
805 dst_strides: &[isize],
806 a_ptr: *const A,
807 a_strides: &[isize],
808 b_ptr: *const B,
809 b_strides: &[isize],
810) -> bool {
811 let Ok(total) = total_len(dims) else {
814 return false;
815 };
816 if total == 0 {
817 return true;
818 }
819 if total <= CONTIGUOUS_RANGE_MIN_LEN {
820 return false;
821 }
822
823 let Some(plan) = contiguous_mul_range_plan(dims, dst_strides, a_strides, b_strides) else {
824 return false;
825 };
826 if !plan.walks_inputs_without_blocking() {
827 return false;
828 }
829
830 let inner_len = plan.inner_len.max(1);
831 let row_len = plan.row_len.max(1);
832 let block_len = inner_len.saturating_mul(row_len).max(1);
833 let outer_groups = total.div_ceil(block_len);
834
835 #[cfg(feature = "parallel")]
836 {
837 let nthreads = crate::execution_policy::rayon_threads();
838 if nthreads > 1 {
839 use crate::threading::{parallel_for_each, SendPtr};
840
841 let dst = SendPtr(dst_ptr);
842 let a = SendPtr(a_ptr as *mut A);
843 let b = SendPtr(b_ptr as *mut B);
844
845 if outer_groups < nthreads {
846 let chunk_len = total.div_ceil(nthreads);
847 let nchunks = total.div_ceil(chunk_len);
848
849 parallel_for_each(0..nchunks, nthreads, &|chunks| {
850 for chunk in chunks {
851 let start = chunk * chunk_len;
852 let end = (start + chunk_len).min(total);
853 let mut index = start;
854
855 while index < end {
856 let in_inner = index % inner_len;
857 let len = (inner_len - in_inner).min(end - index);
858 let a_offset = strided_offset_for_contiguous_linear_index(
859 dims,
860 a_strides,
861 &plan.axis_order,
862 index,
863 );
864 let b_offset = strided_offset_for_contiguous_linear_index(
865 dims,
866 b_strides,
867 &plan.axis_order,
868 index,
869 );
870
871 unsafe {
872 inner_loop_mul2::<O, D, A, B>(
873 dst.as_ptr().add(index),
874 1,
875 a.as_const().offset(a_offset),
876 plan.a_fast_stride,
877 b.as_const().offset(b_offset),
878 plan.b_fast_stride,
879 len,
880 );
881 }
882 index += len;
883 }
884 }
885 });
886
887 return true;
888 }
889
890 let groups_per_chunk = outer_groups.div_ceil(nthreads);
891 let nchunks = outer_groups.div_ceil(groups_per_chunk);
892
893 parallel_for_each(0..nchunks, nthreads, &|chunks| {
894 for chunk in chunks {
895 let group_start = chunk * groups_per_chunk;
896 let group_end = (group_start + groups_per_chunk).min(outer_groups);
897 let mut cursor = ContiguousMulOuterCursor::new(
898 dims,
899 a_strides,
900 b_strides,
901 &plan,
902 group_start,
903 );
904
905 for group in group_start..group_end {
906 let index = group * block_len;
907 unsafe {
908 run_contiguous_mul_row_block::<O, D, A, B>(
909 dst.as_ptr(),
910 a.as_const(),
911 b.as_const(),
912 &plan,
913 index,
914 total,
915 cursor.a_offset,
916 cursor.b_offset,
917 );
918 }
919 cursor.advance();
920 }
921 }
922 });
923
924 true
925 } else {
926 run_contiguous_range_mul_single_thread::<O, D, A, B>(
927 dst_ptr,
928 dims,
929 a_ptr,
930 a_strides,
931 b_ptr,
932 b_strides,
933 &plan,
934 total,
935 block_len,
936 outer_groups,
937 )
938 }
939 }
940
941 #[cfg(not(feature = "parallel"))]
942 {
943 run_contiguous_range_mul_single_thread::<O, D, A, B>(
944 dst_ptr,
945 dims,
946 a_ptr,
947 a_strides,
948 b_ptr,
949 b_strides,
950 &plan,
951 total,
952 block_len,
953 outer_groups,
954 )
955 }
956}
957
958fn run_contiguous_range_mul_single_thread<
959 O: MulOutput<D>,
960 D: Copy + MaybeSendSync + 'static,
961 A: Copy + MaybeSendSync + Mul<B, Output = D> + 'static,
962 B: Copy + MaybeSendSync + 'static,
963>(
964 dst_ptr: *mut O::Slot,
965 dims: &[usize],
966 a_ptr: *const A,
967 a_strides: &[isize],
968 b_ptr: *const B,
969 b_strides: &[isize],
970 plan: &ContiguousMulRangePlan,
971 total: usize,
972 block_len: usize,
973 outer_groups: usize,
974) -> bool {
975 let mut cursor = ContiguousMulOuterCursor::new(dims, a_strides, b_strides, plan, 0);
976 for group in 0..outer_groups {
977 let index = group * block_len;
978 unsafe {
979 run_contiguous_mul_row_block::<O, D, A, B>(
980 dst_ptr,
981 a_ptr,
982 b_ptr,
983 plan,
984 index,
985 total,
986 cursor.a_offset,
987 cursor.b_offset,
988 );
989 }
990 cursor.advance();
991 }
992 true
993}
994
995#[inline(always)]
997unsafe fn inner_loop_map3<
998 D: Copy,
999 A: Copy,
1000 B: Copy,
1001 C: Copy,
1002 OpA: ElementOp<A>,
1003 OpB: ElementOp<B>,
1004 OpC: ElementOp<C>,
1005>(
1006 dp: *mut D,
1007 ds: isize,
1008 ap: *const A,
1009 a_s: isize,
1010 bp: *const B,
1011 b_s: isize,
1012 cp: *const C,
1013 c_s: isize,
1014 len: usize,
1015 f: &impl Fn(A, B, C) -> D,
1016) {
1017 if ds == 1 && a_s == 1 && b_s == 1 && c_s == 1 {
1018 let src_a = std::slice::from_raw_parts(ap, len);
1019 let src_b = std::slice::from_raw_parts(bp, len);
1020 let src_c = std::slice::from_raw_parts(cp, len);
1021 let dst = std::slice::from_raw_parts_mut(dp, len);
1022 simd::dispatch_if_large(len, || {
1023 for (((d, &a), &b), &c) in dst.iter_mut().zip(src_a).zip(src_b).zip(src_c) {
1024 *d = f(OpA::apply(a), OpB::apply(b), OpC::apply(c));
1025 }
1026 });
1027 } else {
1028 let mut dp = dp;
1029 let mut ap = ap;
1030 let mut bp = bp;
1031 let mut cp = cp;
1032 for _ in 0..len {
1033 *dp = f(OpA::apply(*ap), OpB::apply(*bp), OpC::apply(*cp));
1034 dp = dp.offset(ds);
1035 ap = ap.offset(a_s);
1036 bp = bp.offset(b_s);
1037 cp = cp.offset(c_s);
1038 }
1039 }
1040}
1041
1042#[inline(always)]
1044unsafe fn inner_loop_map4<
1045 D: Copy,
1046 A: Copy,
1047 B: Copy,
1048 C: Copy,
1049 E: Copy,
1050 OpA: ElementOp<A>,
1051 OpB: ElementOp<B>,
1052 OpC: ElementOp<C>,
1053 OpE: ElementOp<E>,
1054>(
1055 dp: *mut D,
1056 ds: isize,
1057 ap: *const A,
1058 a_s: isize,
1059 bp: *const B,
1060 b_s: isize,
1061 cp: *const C,
1062 c_s: isize,
1063 ep: *const E,
1064 e_s: isize,
1065 len: usize,
1066 f: &impl Fn(A, B, C, E) -> D,
1067) {
1068 if ds == 1 && a_s == 1 && b_s == 1 && c_s == 1 && e_s == 1 {
1069 let src_a = std::slice::from_raw_parts(ap, len);
1070 let src_b = std::slice::from_raw_parts(bp, len);
1071 let src_c = std::slice::from_raw_parts(cp, len);
1072 let src_e = std::slice::from_raw_parts(ep, len);
1073 let dst = std::slice::from_raw_parts_mut(dp, len);
1074 simd::dispatch_if_large(len, || {
1075 for i in 0..len {
1076 dst[i] = f(
1077 OpA::apply(src_a[i]),
1078 OpB::apply(src_b[i]),
1079 OpC::apply(src_c[i]),
1080 OpE::apply(src_e[i]),
1081 );
1082 }
1083 });
1084 } else {
1085 let mut dp = dp;
1086 let mut ap = ap;
1087 let mut bp = bp;
1088 let mut cp = cp;
1089 let mut ep = ep;
1090 for _ in 0..len {
1091 *dp = f(
1092 OpA::apply(*ap),
1093 OpB::apply(*bp),
1094 OpC::apply(*cp),
1095 OpE::apply(*ep),
1096 );
1097 dp = dp.offset(ds);
1098 ap = ap.offset(a_s);
1099 bp = bp.offset(b_s);
1100 cp = cp.offset(c_s);
1101 ep = ep.offset(e_s);
1102 }
1103 }
1104}
1105
1106pub fn map_into<D: Copy + MaybeSendSync, A: Copy + MaybeSendSync, Op: ElementOp<A>>(
1111 dest: &mut StridedViewMut<D>,
1112 src: &StridedView<A, Op>,
1113 f: impl Fn(A) -> D + MaybeSync,
1114) -> Result<()> {
1115 map_parts_into::<D, A, Op>(
1116 dest.as_mut_ptr(),
1117 dest.dims(),
1118 dest.strides(),
1119 src.ptr(),
1120 src.dims(),
1121 src.strides(),
1122 f,
1123 )
1124}
1125
1126pub(crate) fn map_into_validated<
1127 D: Copy + MaybeSendSync,
1128 A: Copy + MaybeSendSync,
1129 Op: ElementOp<A>,
1130>(
1131 dest: &mut StridedViewMut<D>,
1132 src: &StridedView<A, Op>,
1133 f: impl Fn(A) -> D + MaybeSync,
1134 validated: ValidatedDestinationLayout,
1135) -> Result<()> {
1136 ensure_same_shape(dest.dims(), src.dims())?;
1137 map_parts_into_validated::<D, A, Op>(
1138 dest.as_mut_ptr(),
1139 dest.dims(),
1140 dest.strides(),
1141 src.ptr(),
1142 src.strides(),
1143 f,
1144 validated,
1145 )
1146}
1147
1148pub(crate) fn map_raw_into<D: Copy + MaybeSendSync, A: Copy + MaybeSendSync, Op: ElementOp<A>>(
1149 dest: &mut crate::RawStridedMut<'_, D>,
1150 src: &crate::RawStridedRef<'_, A>,
1151 f: impl Fn(A) -> D + MaybeSync,
1152) -> Result<()> {
1153 map_parts_into::<D, A, Op>(
1154 dest.as_mut_ptr(),
1155 dest.dims(),
1156 dest.strides(),
1157 src.ptr(),
1158 src.dims(),
1159 src.strides(),
1160 f,
1161 )
1162}
1163
1164pub(crate) fn map_raw_into_validated<
1165 D: Copy + MaybeSendSync,
1166 A: Copy + MaybeSendSync,
1167 Op: ElementOp<A>,
1168>(
1169 dest: &mut crate::RawStridedMut<'_, D>,
1170 src: &crate::RawStridedRef<'_, A>,
1171 f: impl Fn(A) -> D + MaybeSync,
1172 validated: ValidatedDestinationLayout,
1173) -> Result<()> {
1174 ensure_same_shape(dest.dims(), src.dims())?;
1175 map_parts_into_validated::<D, A, Op>(
1176 dest.as_mut_ptr(),
1177 dest.dims(),
1178 dest.strides(),
1179 src.ptr(),
1180 src.strides(),
1181 f,
1182 validated,
1183 )
1184}
1185
1186#[allow(clippy::too_many_arguments)]
1187fn map_parts_into<D: Copy + MaybeSendSync, A: Copy + MaybeSendSync, Op: ElementOp<A>>(
1188 dst_ptr: *mut D,
1189 dst_dims: &[usize],
1190 dst_strides: &[isize],
1191 src_ptr: *const A,
1192 src_dims: &[usize],
1193 src_strides: &[isize],
1194 f: impl Fn(A) -> D + MaybeSync,
1195) -> Result<()> {
1196 ensure_same_shape(dst_dims, src_dims)?;
1197 let validated = validate_destination_layout(dst_dims, dst_strides)?;
1198 map_parts_into_validated::<D, A, Op>(
1199 dst_ptr,
1200 dst_dims,
1201 dst_strides,
1202 src_ptr,
1203 src_strides,
1204 f,
1205 validated,
1206 )
1207}
1208
1209#[allow(clippy::too_many_arguments)]
1210fn map_parts_into_validated<D: Copy + MaybeSendSync, A: Copy + MaybeSendSync, Op: ElementOp<A>>(
1211 dst_ptr: *mut D,
1212 dst_dims: &[usize],
1213 dst_strides: &[isize],
1214 src_ptr: *const A,
1215 src_strides: &[isize],
1216 f: impl Fn(A) -> D + MaybeSync,
1217 _validated: ValidatedDestinationLayout,
1218) -> Result<()> {
1219 if sequential_contiguous_layout(dst_dims, &[dst_strides, src_strides])?.is_some() {
1220 let len = total_len(dst_dims)?;
1221 let dst = unsafe { std::slice::from_raw_parts_mut(dst_ptr, len) };
1222 let src = unsafe { std::slice::from_raw_parts(src_ptr, len) };
1223 simd::dispatch_if_large(len, || {
1224 for i in 0..len {
1225 dst[i] = f(Op::apply(src[i]));
1226 }
1227 });
1228 return Ok(());
1229 }
1230
1231 let strides_list: [&[isize]; 2] = [dst_strides, src_strides];
1232 let elem_size = std::mem::size_of::<D>().max(std::mem::size_of::<A>());
1233 let total = total_len(dst_dims)?;
1234
1235 let (fused_dims, ordered_strides, plan) = if total <= SMALL_TENSOR_THRESHOLD {
1237 build_plan_fused_small(dst_dims, &strides_list)
1238 } else {
1239 build_plan_fused(dst_dims, &strides_list, Some(0), elem_size)
1240 };
1241
1242 #[cfg(feature = "parallel")]
1243 {
1244 let total = total_len(&fused_dims)?;
1245 let nthreads = crate::execution_policy::rayon_threads();
1246 if total > MINTHREADLENGTH && nthreads > 1 {
1247 use crate::threading::SendPtr;
1248 let dst_send = SendPtr(dst_ptr);
1249 let src_send = SendPtr(src_ptr as *mut A);
1250
1251 let costs = compute_costs(&ordered_strides);
1252 let initial_offsets = vec![0isize; strides_list.len()];
1253 return mapreduce_threaded(
1254 &fused_dims,
1255 &plan.block,
1256 &ordered_strides,
1257 &initial_offsets,
1258 &costs,
1259 nthreads,
1260 0,
1261 1,
1262 &|dims, blocks, strides_list, offsets| {
1263 for_each_inner_block_with_offsets(
1264 dims,
1265 blocks,
1266 strides_list,
1267 offsets,
1268 |offsets, len, strides| {
1269 let dp = unsafe { dst_send.as_ptr().offset(offsets[0]) };
1270 let sp = unsafe { src_send.as_const().offset(offsets[1]) };
1271 unsafe {
1272 inner_loop_map1::<D, A, Op>(dp, strides[0], sp, strides[1], len, &f)
1273 };
1274 Ok(())
1275 },
1276 )
1277 },
1278 );
1279 }
1280 }
1281
1282 let initial_offsets = vec![0isize; ordered_strides.len()];
1283 for_each_inner_block_preordered(
1284 &fused_dims,
1285 &plan.block,
1286 &ordered_strides,
1287 &initial_offsets,
1288 |offsets, len, strides| {
1289 let dp = unsafe { dst_ptr.offset(offsets[0]) };
1290 let sp = unsafe { src_ptr.offset(offsets[1]) };
1291 unsafe { inner_loop_map1::<D, A, Op>(dp, strides[0], sp, strides[1], len, &f) };
1292 Ok(())
1293 },
1294 )
1295}
1296
1297pub fn zip_map2_into<
1302 D: Copy + MaybeSendSync,
1303 A: Copy + MaybeSendSync,
1304 B: Copy + MaybeSendSync,
1305 OpA: ElementOp<A>,
1306 OpB: ElementOp<B>,
1307>(
1308 dest: &mut StridedViewMut<D>,
1309 a: &StridedView<A, OpA>,
1310 b: &StridedView<B, OpB>,
1311 f: impl Fn(A, B) -> D + MaybeSync,
1312) -> Result<()> {
1313 zip_map2_parts_into::<D, A, B, OpA, OpB>(
1314 dest.as_mut_ptr(),
1315 dest.dims(),
1316 dest.strides(),
1317 a.ptr(),
1318 a.dims(),
1319 a.strides(),
1320 b.ptr(),
1321 b.dims(),
1322 b.strides(),
1323 f,
1324 )
1325}
1326
1327pub(crate) fn zip_map2_into_validated<
1328 D: Copy + MaybeSendSync,
1329 A: Copy + MaybeSendSync,
1330 B: Copy + MaybeSendSync,
1331 OpA: ElementOp<A>,
1332 OpB: ElementOp<B>,
1333>(
1334 dest: &mut StridedViewMut<D>,
1335 a: &StridedView<A, OpA>,
1336 b: &StridedView<B, OpB>,
1337 f: impl Fn(A, B) -> D + MaybeSync,
1338 validated: ValidatedDestinationLayout,
1339) -> Result<()> {
1340 zip_map2_parts_into_validated::<D, A, B, OpA, OpB>(
1341 dest.as_mut_ptr(),
1342 dest.dims(),
1343 dest.strides(),
1344 a.ptr(),
1345 a.strides(),
1346 b.ptr(),
1347 b.strides(),
1348 f,
1349 validated,
1350 )
1351}
1352
1353#[non_exhaustive]
1355#[derive(Clone, Copy, Debug, Eq, PartialEq)]
1356pub enum CompareOp {
1357 Eq,
1358 Lt,
1359 Le,
1360 Gt,
1361 Ge,
1362}
1363
1364pub fn compare_into<T, OpA, OpB>(
1375 dest: &mut StridedViewMut<bool>,
1376 a: &StridedView<T, OpA>,
1377 b: &StridedView<T, OpB>,
1378 op: CompareOp,
1379) -> Result<()>
1380where
1381 T: Copy + MaybeSendSync + PartialOrd,
1382 OpA: ElementOp<T>,
1383 OpB: ElementOp<T>,
1384{
1385 match op {
1386 CompareOp::Eq => zip_map2_into(dest, a, b, |lhs, rhs| lhs == rhs),
1387 CompareOp::Lt => zip_map2_into(dest, a, b, |lhs, rhs| lhs < rhs),
1388 CompareOp::Le => zip_map2_into(dest, a, b, |lhs, rhs| lhs <= rhs),
1389 CompareOp::Gt => zip_map2_into(dest, a, b, |lhs, rhs| lhs > rhs),
1390 CompareOp::Ge => zip_map2_into(dest, a, b, |lhs, rhs| lhs >= rhs),
1391 }
1392}
1393
1394pub fn compare_into_uninit<T, OpA, OpB>(
1413 dest: &mut StridedViewMut<MaybeUninit<bool>>,
1414 a: &StridedView<T, OpA>,
1415 b: &StridedView<T, OpB>,
1416 op: CompareOp,
1417) -> Result<()>
1418where
1419 T: Copy + MaybeSendSync + PartialOrd,
1420 OpA: ElementOp<T>,
1421 OpB: ElementOp<T>,
1422{
1423 ensure_same_shape(dest.dims(), a.dims())?;
1424 ensure_same_shape(dest.dims(), b.dims())?;
1425 let validated = validate_destination_layout(dest.dims(), dest.strides())?;
1426 validate_typed_no_overlap(dest, a, 0)?;
1427 validate_typed_no_overlap(dest, b, 1)?;
1428 match op {
1429 CompareOp::Eq => zip_map2_into_validated(
1430 dest,
1431 a,
1432 b,
1433 |lhs, rhs| MaybeUninit::new(lhs == rhs),
1434 validated,
1435 ),
1436 CompareOp::Lt => zip_map2_into_validated(
1437 dest,
1438 a,
1439 b,
1440 |lhs, rhs| MaybeUninit::new(lhs < rhs),
1441 validated,
1442 ),
1443 CompareOp::Le => zip_map2_into_validated(
1444 dest,
1445 a,
1446 b,
1447 |lhs, rhs| MaybeUninit::new(lhs <= rhs),
1448 validated,
1449 ),
1450 CompareOp::Gt => zip_map2_into_validated(
1451 dest,
1452 a,
1453 b,
1454 |lhs, rhs| MaybeUninit::new(lhs > rhs),
1455 validated,
1456 ),
1457 CompareOp::Ge => zip_map2_into_validated(
1458 dest,
1459 a,
1460 b,
1461 |lhs, rhs| MaybeUninit::new(lhs >= rhs),
1462 validated,
1463 ),
1464 }
1465}
1466
1467pub(crate) fn zip_map2_raw_into_validated<
1468 D: Copy + MaybeSendSync,
1469 A: Copy + MaybeSendSync,
1470 B: Copy + MaybeSendSync,
1471 OpA: ElementOp<A>,
1472 OpB: ElementOp<B>,
1473>(
1474 dest: &mut crate::RawStridedMut<'_, D>,
1475 a: &crate::RawStridedRef<'_, A>,
1476 b: &crate::RawStridedRef<'_, B>,
1477 f: impl Fn(A, B) -> D + MaybeSync,
1478 validated: ValidatedDestinationLayout,
1479) -> Result<()> {
1480 ensure_same_shape(dest.dims(), a.dims())?;
1481 ensure_same_shape(dest.dims(), b.dims())?;
1482 zip_map2_parts_into_validated::<D, A, B, OpA, OpB>(
1483 dest.as_mut_ptr(),
1484 dest.dims(),
1485 dest.strides(),
1486 a.ptr(),
1487 a.strides(),
1488 b.ptr(),
1489 b.strides(),
1490 f,
1491 validated,
1492 )
1493}
1494
1495#[allow(clippy::too_many_arguments)]
1496fn zip_map2_parts_into<
1497 D: Copy + MaybeSendSync,
1498 A: Copy + MaybeSendSync,
1499 B: Copy + MaybeSendSync,
1500 OpA: ElementOp<A>,
1501 OpB: ElementOp<B>,
1502>(
1503 dst_ptr: *mut D,
1504 dst_dims: &[usize],
1505 dst_strides: &[isize],
1506 a_ptr: *const A,
1507 a_dims: &[usize],
1508 a_strides: &[isize],
1509 b_ptr: *const B,
1510 b_dims: &[usize],
1511 b_strides: &[isize],
1512 f: impl Fn(A, B) -> D + MaybeSync,
1513) -> Result<()> {
1514 ensure_same_shape(dst_dims, a_dims)?;
1515 ensure_same_shape(dst_dims, b_dims)?;
1516 let validated = validate_destination_layout(dst_dims, dst_strides)?;
1517 zip_map2_parts_into_validated::<D, A, B, OpA, OpB>(
1518 dst_ptr,
1519 dst_dims,
1520 dst_strides,
1521 a_ptr,
1522 a_strides,
1523 b_ptr,
1524 b_strides,
1525 f,
1526 validated,
1527 )
1528}
1529
1530#[allow(clippy::too_many_arguments)]
1531fn zip_map2_parts_into_validated<
1532 D: Copy + MaybeSendSync,
1533 A: Copy + MaybeSendSync,
1534 B: Copy + MaybeSendSync,
1535 OpA: ElementOp<A>,
1536 OpB: ElementOp<B>,
1537>(
1538 dst_ptr: *mut D,
1539 dst_dims: &[usize],
1540 dst_strides: &[isize],
1541 a_ptr: *const A,
1542 a_strides: &[isize],
1543 b_ptr: *const B,
1544 b_strides: &[isize],
1545 f: impl Fn(A, B) -> D + MaybeSync,
1546 _validated: ValidatedDestinationLayout,
1547) -> Result<()> {
1548 if sequential_contiguous_layout(dst_dims, &[dst_strides, a_strides, b_strides])?.is_some() {
1549 let len = total_len(dst_dims)?;
1550 let dst = unsafe { std::slice::from_raw_parts_mut(dst_ptr, len) };
1551 let sa = unsafe { std::slice::from_raw_parts(a_ptr, len) };
1552 let sb = unsafe { std::slice::from_raw_parts(b_ptr, len) };
1553 simd::dispatch_if_large(len, || {
1554 for i in 0..len {
1555 dst[i] = f(OpA::apply(sa[i]), OpB::apply(sb[i]));
1556 }
1557 });
1558 return Ok(());
1559 }
1560
1561 let strides_list: [&[isize]; 3] = [dst_strides, a_strides, b_strides];
1562 let elem_size = std::mem::size_of::<D>()
1563 .max(std::mem::size_of::<A>())
1564 .max(std::mem::size_of::<B>());
1565 let total = total_len(dst_dims)?;
1566
1567 let (fused_dims, ordered_strides, plan) = if total <= SMALL_TENSOR_THRESHOLD {
1569 build_plan_fused_small(dst_dims, &strides_list)
1570 } else {
1571 build_plan_fused(dst_dims, &strides_list, Some(0), elem_size)
1572 };
1573
1574 #[cfg(feature = "parallel")]
1575 {
1576 let total = total_len(&fused_dims)?;
1577 let nthreads = crate::execution_policy::rayon_threads();
1578 if total > MINTHREADLENGTH && nthreads > 1 {
1579 use crate::threading::SendPtr;
1580 let dst_send = SendPtr(dst_ptr);
1581 let a_send = SendPtr(a_ptr as *mut A);
1582 let b_send = SendPtr(b_ptr as *mut B);
1583
1584 let costs = compute_costs(&ordered_strides);
1585 let initial_offsets = vec![0isize; strides_list.len()];
1586 return mapreduce_threaded(
1587 &fused_dims,
1588 &plan.block,
1589 &ordered_strides,
1590 &initial_offsets,
1591 &costs,
1592 nthreads,
1593 0,
1594 1,
1595 &|dims, blocks, strides_list, offsets| {
1596 for_each_inner_block_with_offsets(
1597 dims,
1598 blocks,
1599 strides_list,
1600 offsets,
1601 |offsets, len, strides| {
1602 let dp = unsafe { dst_send.as_ptr().offset(offsets[0]) };
1603 let ap = unsafe { a_send.as_const().offset(offsets[1]) };
1604 let bp = unsafe { b_send.as_const().offset(offsets[2]) };
1605 unsafe {
1606 inner_loop_map2::<D, A, B, OpA, OpB>(
1607 dp, strides[0], ap, strides[1], bp, strides[2], len, &f,
1608 )
1609 };
1610 Ok(())
1611 },
1612 )
1613 },
1614 );
1615 }
1616 }
1617
1618 let initial_offsets = vec![0isize; ordered_strides.len()];
1619 for_each_inner_block_preordered(
1620 &fused_dims,
1621 &plan.block,
1622 &ordered_strides,
1623 &initial_offsets,
1624 |offsets, len, strides| {
1625 let dp = unsafe { dst_ptr.offset(offsets[0]) };
1626 let ap = unsafe { a_ptr.offset(offsets[1]) };
1627 let bp = unsafe { b_ptr.offset(offsets[2]) };
1628 unsafe {
1629 inner_loop_map2::<D, A, B, OpA, OpB>(
1630 dp, strides[0], ap, strides[1], bp, strides[2], len, &f,
1631 )
1632 };
1633 Ok(())
1634 },
1635 )
1636}
1637
1638fn mul_identity_into_raw<
1639 O: MulOutput<D>,
1640 D: Copy + MaybeSendSync + 'static,
1641 A: Copy + MaybeSendSync + Mul<B, Output = D> + 'static,
1642 B: Copy + MaybeSendSync + 'static,
1643>(
1644 dst_ptr: *mut O::Slot,
1645 dst_dims: &[usize],
1646 dst_strides: &[isize],
1647 a_ptr: *const A,
1648 a_strides: &[isize],
1649 b_ptr: *const B,
1650 b_strides: &[isize],
1651 _validated: ValidatedDestinationLayout,
1652) -> Result<()> {
1653 debug_assert_eq!(dst_dims.len(), a_strides.len());
1654 debug_assert_eq!(dst_dims.len(), b_strides.len());
1655
1656 if sequential_contiguous_layout(dst_dims, &[dst_strides, a_strides, b_strides])?.is_some() {
1657 let len = total_len(dst_dims)?;
1658 let sa = unsafe { std::slice::from_raw_parts(a_ptr, len) };
1659 let sb = unsafe { std::slice::from_raw_parts(b_ptr, len) };
1660 if unsafe { O::try_contiguous(dst_ptr, len, sa, sb) } {
1661 return Ok(());
1662 }
1663 for i in 0..len {
1664 unsafe { O::write(dst_ptr.add(i), multiply_value(sa[i], sb[i])) };
1665 }
1666 return Ok(());
1667 }
1668
1669 let strides_list: [&[isize]; 3] = [dst_strides, a_strides, b_strides];
1670 let elem_size = std::mem::size_of::<D>()
1671 .max(std::mem::size_of::<A>())
1672 .max(std::mem::size_of::<B>());
1673 let total = total_len(dst_dims)?;
1674
1675 if try_contiguous_range_mul::<O, D, A, B>(
1676 dst_ptr,
1677 dst_dims,
1678 dst_strides,
1679 a_ptr,
1680 a_strides,
1681 b_ptr,
1682 b_strides,
1683 ) {
1684 return Ok(());
1685 }
1686
1687 let (fused_dims, ordered_strides, plan) = if total <= SMALL_TENSOR_THRESHOLD {
1688 build_plan_fused_small(dst_dims, &strides_list)
1689 } else {
1690 build_plan_fused(dst_dims, &strides_list, Some(0), elem_size)
1691 };
1692
1693 #[cfg(feature = "parallel")]
1694 {
1695 let total = total_len(&fused_dims)?;
1696 let nthreads = crate::execution_policy::rayon_threads();
1697 if total > MINTHREADLENGTH && nthreads > 1 {
1698 use crate::threading::SendPtr;
1699 let dst_send = SendPtr(dst_ptr);
1700 let a_send = SendPtr(a_ptr as *mut A);
1701 let b_send = SendPtr(b_ptr as *mut B);
1702
1703 let costs = compute_costs(&ordered_strides);
1704 let initial_offsets = vec![0isize; strides_list.len()];
1705 return mapreduce_threaded(
1706 &fused_dims,
1707 &plan.block,
1708 &ordered_strides,
1709 &initial_offsets,
1710 &costs,
1711 nthreads,
1712 0,
1713 1,
1714 &|dims, blocks, strides_list, offsets| {
1715 for_each_inner_block_with_offsets(
1716 dims,
1717 blocks,
1718 strides_list,
1719 offsets,
1720 |offsets, len, strides| {
1721 let dp = unsafe { dst_send.as_ptr().offset(offsets[0]) };
1722 let ap = unsafe { a_send.as_const().offset(offsets[1]) };
1723 let bp = unsafe { b_send.as_const().offset(offsets[2]) };
1724 unsafe {
1725 inner_loop_mul2::<O, D, A, B>(
1726 dp, strides[0], ap, strides[1], bp, strides[2], len,
1727 )
1728 };
1729 Ok(())
1730 },
1731 )
1732 },
1733 );
1734 }
1735 }
1736
1737 let initial_offsets = vec![0isize; ordered_strides.len()];
1738 for_each_inner_block_preordered(
1739 &fused_dims,
1740 &plan.block,
1741 &ordered_strides,
1742 &initial_offsets,
1743 |offsets, len, strides| {
1744 let dp = unsafe { dst_ptr.offset(offsets[0]) };
1745 let ap = unsafe { a_ptr.offset(offsets[1]) };
1746 let bp = unsafe { b_ptr.offset(offsets[2]) };
1747 unsafe {
1748 inner_loop_mul2::<O, D, A, B>(dp, strides[0], ap, strides[1], bp, strides[2], len)
1749 };
1750 Ok(())
1751 },
1752 )
1753}
1754
1755pub fn mul_into<
1760 D: Copy + MaybeSendSync + 'static,
1761 A: Copy + Mul<B, Output = D> + MaybeSendSync + 'static,
1762 B: Copy + MaybeSendSync + 'static,
1763 OpA: ElementOp<A>,
1764 OpB: ElementOp<B>,
1765>(
1766 dest: &mut StridedViewMut<D>,
1767 a: &StridedView<A, OpA>,
1768 b: &StridedView<B, OpB>,
1769) -> Result<()> {
1770 ensure_same_shape(dest.dims(), a.dims())?;
1771 ensure_same_shape(dest.dims(), b.dims())?;
1772 let validated = validate_destination_layout(dest.dims(), dest.strides())?;
1773
1774 if OpA::IS_IDENTITY && OpB::IS_IDENTITY {
1775 return mul_identity_into_raw::<InitializedOutput, D, A, B>(
1776 dest.as_mut_ptr(),
1777 dest.dims(),
1778 dest.strides(),
1779 a.ptr(),
1780 a.strides(),
1781 b.ptr(),
1782 b.strides(),
1783 validated,
1784 );
1785 }
1786
1787 zip_map2_into_validated(dest, a, b, multiply_value, validated)
1788}
1789
1790pub fn mul_into_uninit<
1809 D: Copy + MaybeSendSync + 'static,
1810 A: Copy + Mul<B, Output = D> + MaybeSendSync + 'static,
1811 B: Copy + MaybeSendSync + 'static,
1812 OpA: ElementOp<A>,
1813 OpB: ElementOp<B>,
1814>(
1815 dest: &mut StridedViewMut<MaybeUninit<D>>,
1816 a: &StridedView<A, OpA>,
1817 b: &StridedView<B, OpB>,
1818) -> Result<()> {
1819 ensure_same_shape(dest.dims(), a.dims())?;
1820 ensure_same_shape(dest.dims(), b.dims())?;
1821 let validated = validate_destination_layout(dest.dims(), dest.strides())?;
1822 validate_typed_no_overlap(dest, a, 0)?;
1823 validate_typed_no_overlap(dest, b, 1)?;
1824
1825 if OpA::IS_IDENTITY && OpB::IS_IDENTITY {
1826 return mul_identity_into_raw::<UninitializedOutput, D, A, B>(
1827 dest.as_mut_ptr(),
1828 dest.dims(),
1829 dest.strides(),
1830 a.ptr(),
1831 a.strides(),
1832 b.ptr(),
1833 b.strides(),
1834 validated,
1835 );
1836 }
1837
1838 zip_map2_into_validated(
1839 dest,
1840 a,
1841 b,
1842 |lhs, rhs| MaybeUninit::new(multiply_value(lhs, rhs)),
1843 validated,
1844 )
1845}
1846
1847fn broadcast_strides_for_axes(
1848 source_dims: &[usize],
1849 source_strides: &[isize],
1850 target_dims: &[usize],
1851 axes: &[usize],
1852) -> Result<AxisVec<isize>> {
1853 if source_dims.len() != axes.len() {
1854 return Err(StridedError::RankMismatch(source_dims.len(), axes.len()));
1855 }
1856 debug_assert_eq!(source_dims.len(), source_strides.len());
1857
1858 let mut seen = AxisVec::<bool>::new();
1859 seen.resize(target_dims.len(), false);
1860 let mut strides = AxisVec::<isize>::new();
1861 strides.resize(target_dims.len(), 0);
1862 for (src_axis, &dst_axis) in axes.iter().enumerate() {
1863 if dst_axis >= target_dims.len() {
1864 return Err(StridedError::InvalidAxis {
1865 axis: dst_axis,
1866 rank: target_dims.len(),
1867 });
1868 }
1869 if seen[dst_axis] {
1870 return Err(StridedError::InvalidAxis {
1871 axis: dst_axis,
1872 rank: target_dims.len(),
1873 });
1874 }
1875 seen[dst_axis] = true;
1876
1877 let source_dim = source_dims[src_axis];
1878 let target_dim = target_dims[dst_axis];
1879 if source_dim != target_dim && source_dim != 1 {
1880 return Err(StridedError::ShapeMismatch(
1881 source_dims.to_vec(),
1882 target_dims.to_vec(),
1883 ));
1884 }
1885 if source_dim == target_dim {
1886 strides[dst_axis] = source_strides[src_axis];
1887 }
1888 }
1889
1890 Ok(strides)
1891}
1892
1893fn broadcast_view_with_strides<'a, T, Op: ElementOp<T>>(
1894 view: &StridedView<'a, T, Op>,
1895 target_dims: &[usize],
1896 strides: &[isize],
1897) -> StridedView<'a, T, Op> {
1898 unsafe { StridedView::new_unchecked(view.data(), target_dims, strides, view.offset()) }
1899}
1900
1901pub fn broadcast_mul_into<
1906 D: Copy + MaybeSendSync + 'static,
1907 A: Copy + Mul<B, Output = D> + MaybeSendSync + 'static,
1908 B: Copy + MaybeSendSync + 'static,
1909 OpA: ElementOp<A>,
1910 OpB: ElementOp<B>,
1911>(
1912 dest: &mut StridedViewMut<D>,
1913 a: &StridedView<A, OpA>,
1914 a_axes: &[usize],
1915 b: &StridedView<B, OpB>,
1916 b_axes: &[usize],
1917) -> Result<()> {
1918 let validated = validate_destination_layout(dest.dims(), dest.strides())?;
1919 let a_strides = broadcast_strides_for_axes(a.dims(), a.strides(), dest.dims(), a_axes)?;
1920 let b_strides = broadcast_strides_for_axes(b.dims(), b.strides(), dest.dims(), b_axes)?;
1921
1922 if OpA::IS_IDENTITY && OpB::IS_IDENTITY {
1923 return mul_identity_into_raw::<InitializedOutput, D, A, B>(
1924 dest.as_mut_ptr(),
1925 dest.dims(),
1926 dest.strides(),
1927 a.ptr(),
1928 &a_strides,
1929 b.ptr(),
1930 &b_strides,
1931 validated,
1932 );
1933 }
1934
1935 let a = broadcast_view_with_strides(a, dest.dims(), &a_strides);
1936 let b = broadcast_view_with_strides(b, dest.dims(), &b_strides);
1937 zip_map2_into_validated(dest, &a, &b, multiply_value, validated)
1938}
1939
1940pub fn broadcast_mul_into_uninit<
1959 D: Copy + MaybeSendSync + 'static,
1960 A: Copy + Mul<B, Output = D> + MaybeSendSync + 'static,
1961 B: Copy + MaybeSendSync + 'static,
1962 OpA: ElementOp<A>,
1963 OpB: ElementOp<B>,
1964>(
1965 dest: &mut StridedViewMut<MaybeUninit<D>>,
1966 a: &StridedView<A, OpA>,
1967 a_axes: &[usize],
1968 b: &StridedView<B, OpB>,
1969 b_axes: &[usize],
1970) -> Result<()> {
1971 let validated = validate_destination_layout(dest.dims(), dest.strides())?;
1972 let a_strides = broadcast_strides_for_axes(a.dims(), a.strides(), dest.dims(), a_axes)?;
1973 let b_strides = broadcast_strides_for_axes(b.dims(), b.strides(), dest.dims(), b_axes)?;
1974 validate_typed_no_overlap(dest, a, 0)?;
1975 validate_typed_no_overlap(dest, b, 1)?;
1976
1977 if OpA::IS_IDENTITY && OpB::IS_IDENTITY {
1978 return mul_identity_into_raw::<UninitializedOutput, D, A, B>(
1979 dest.as_mut_ptr(),
1980 dest.dims(),
1981 dest.strides(),
1982 a.ptr(),
1983 &a_strides,
1984 b.ptr(),
1985 &b_strides,
1986 validated,
1987 );
1988 }
1989
1990 let a = broadcast_view_with_strides(a, dest.dims(), &a_strides);
1991 let b = broadcast_view_with_strides(b, dest.dims(), &b_strides);
1992 zip_map2_into_validated(
1993 dest,
1994 &a,
1995 &b,
1996 |lhs, rhs| MaybeUninit::new(multiply_value(lhs, rhs)),
1997 validated,
1998 )
1999}
2000
2001pub fn zip_map3_into<
2003 D: Copy + MaybeSendSync,
2004 A: Copy + MaybeSendSync,
2005 B: Copy + MaybeSendSync,
2006 C: Copy + MaybeSendSync,
2007 OpA: ElementOp<A>,
2008 OpB: ElementOp<B>,
2009 OpC: ElementOp<C>,
2010>(
2011 dest: &mut StridedViewMut<D>,
2012 a: &StridedView<A, OpA>,
2013 b: &StridedView<B, OpB>,
2014 c: &StridedView<C, OpC>,
2015 f: impl Fn(A, B, C) -> D + MaybeSync,
2016) -> Result<()> {
2017 ensure_same_shape(dest.dims(), a.dims())?;
2018 ensure_same_shape(dest.dims(), b.dims())?;
2019 ensure_same_shape(dest.dims(), c.dims())?;
2020 let validated = validate_destination_layout(dest.dims(), dest.strides())?;
2021 zip_map3_into_validated(dest, a, b, c, f, validated)
2022}
2023
2024pub(crate) fn zip_map3_into_validated<
2025 D: Copy + MaybeSendSync,
2026 A: Copy + MaybeSendSync,
2027 B: Copy + MaybeSendSync,
2028 C: Copy + MaybeSendSync,
2029 OpA: ElementOp<A>,
2030 OpB: ElementOp<B>,
2031 OpC: ElementOp<C>,
2032>(
2033 dest: &mut StridedViewMut<D>,
2034 a: &StridedView<A, OpA>,
2035 b: &StridedView<B, OpB>,
2036 c: &StridedView<C, OpC>,
2037 f: impl Fn(A, B, C) -> D + MaybeSync,
2038 _validated: ValidatedDestinationLayout,
2039) -> Result<()> {
2040 ensure_same_shape(dest.dims(), a.dims())?;
2041 ensure_same_shape(dest.dims(), b.dims())?;
2042 ensure_same_shape(dest.dims(), c.dims())?;
2043 let dst_ptr = dest.as_mut_ptr();
2044 let a_ptr = a.ptr();
2045 let b_ptr = b.ptr();
2046 let c_ptr = c.ptr();
2047
2048 let dst_dims = dest.dims();
2049 let dst_strides = dest.strides();
2050
2051 if sequential_contiguous_layout(
2052 dst_dims,
2053 &[dst_strides, a.strides(), b.strides(), c.strides()],
2054 )?
2055 .is_some()
2056 {
2057 let len = total_len(dst_dims)?;
2058 let dst = unsafe { std::slice::from_raw_parts_mut(dst_ptr, len) };
2059 let sa = unsafe { std::slice::from_raw_parts(a_ptr, len) };
2060 let sb = unsafe { std::slice::from_raw_parts(b_ptr, len) };
2061 let sc = unsafe { std::slice::from_raw_parts(c_ptr, len) };
2062 simd::dispatch_if_large(len, || {
2063 for (((d, &a), &b), &c) in dst.iter_mut().zip(sa).zip(sb).zip(sc) {
2064 *d = f(OpA::apply(a), OpB::apply(b), OpC::apply(c));
2065 }
2066 });
2067 return Ok(());
2068 }
2069
2070 let strides_list: [&[isize]; 4] = [dst_strides, a.strides(), b.strides(), c.strides()];
2071 let elem_size = std::mem::size_of::<D>()
2072 .max(std::mem::size_of::<A>())
2073 .max(std::mem::size_of::<B>())
2074 .max(std::mem::size_of::<C>());
2075 let total = total_len(dst_dims)?;
2076
2077 let (fused_dims, ordered_strides, plan) = if total <= SMALL_TENSOR_THRESHOLD {
2079 build_plan_fused_small(dst_dims, &strides_list)
2080 } else {
2081 build_plan_fused(dst_dims, &strides_list, Some(0), elem_size)
2082 };
2083
2084 #[cfg(feature = "parallel")]
2085 {
2086 let total = total_len(&fused_dims)?;
2087 let nthreads = crate::execution_policy::rayon_threads();
2088 if total > MINTHREADLENGTH && nthreads > 1 {
2089 use crate::threading::SendPtr;
2090 let dst_send = SendPtr(dst_ptr);
2091 let a_send = SendPtr(a_ptr as *mut A);
2092 let b_send = SendPtr(b_ptr as *mut B);
2093 let c_send = SendPtr(c_ptr as *mut C);
2094
2095 let costs = compute_costs(&ordered_strides);
2096 let initial_offsets = vec![0isize; strides_list.len()];
2097 return mapreduce_threaded(
2098 &fused_dims,
2099 &plan.block,
2100 &ordered_strides,
2101 &initial_offsets,
2102 &costs,
2103 nthreads,
2104 0,
2105 1,
2106 &|dims, blocks, strides_list, offsets| {
2107 for_each_inner_block_with_offsets(
2108 dims,
2109 blocks,
2110 strides_list,
2111 offsets,
2112 |offsets, len, strides| {
2113 let dp = unsafe { dst_send.as_ptr().offset(offsets[0]) };
2114 let ap = unsafe { a_send.as_const().offset(offsets[1]) };
2115 let bp = unsafe { b_send.as_const().offset(offsets[2]) };
2116 let cp = unsafe { c_send.as_const().offset(offsets[3]) };
2117 unsafe {
2118 inner_loop_map3::<D, A, B, C, OpA, OpB, OpC>(
2119 dp, strides[0], ap, strides[1], bp, strides[2], cp, strides[3],
2120 len, &f,
2121 )
2122 };
2123 Ok(())
2124 },
2125 )
2126 },
2127 );
2128 }
2129 }
2130
2131 let initial_offsets = vec![0isize; ordered_strides.len()];
2132 for_each_inner_block_preordered(
2133 &fused_dims,
2134 &plan.block,
2135 &ordered_strides,
2136 &initial_offsets,
2137 |offsets, len, strides| {
2138 let dp = unsafe { dst_ptr.offset(offsets[0]) };
2139 let ap = unsafe { a_ptr.offset(offsets[1]) };
2140 let bp = unsafe { b_ptr.offset(offsets[2]) };
2141 let cp = unsafe { c_ptr.offset(offsets[3]) };
2142 unsafe {
2143 inner_loop_map3::<D, A, B, C, OpA, OpB, OpC>(
2144 dp, strides[0], ap, strides[1], bp, strides[2], cp, strides[3], len, &f,
2145 )
2146 };
2147 Ok(())
2148 },
2149 )
2150}
2151
2152pub fn zip_map4_into<
2154 D: Copy + MaybeSendSync,
2155 A: Copy + MaybeSendSync,
2156 B: Copy + MaybeSendSync,
2157 C: Copy + MaybeSendSync,
2158 E: Copy + MaybeSendSync,
2159 OpA: ElementOp<A>,
2160 OpB: ElementOp<B>,
2161 OpC: ElementOp<C>,
2162 OpE: ElementOp<E>,
2163>(
2164 dest: &mut StridedViewMut<D>,
2165 a: &StridedView<A, OpA>,
2166 b: &StridedView<B, OpB>,
2167 c: &StridedView<C, OpC>,
2168 e: &StridedView<E, OpE>,
2169 f: impl Fn(A, B, C, E) -> D + MaybeSync,
2170) -> Result<()> {
2171 ensure_same_shape(dest.dims(), a.dims())?;
2172 ensure_same_shape(dest.dims(), b.dims())?;
2173 ensure_same_shape(dest.dims(), c.dims())?;
2174 ensure_same_shape(dest.dims(), e.dims())?;
2175 let validated = validate_destination_layout(dest.dims(), dest.strides())?;
2176 zip_map4_into_validated(dest, a, b, c, e, f, validated)
2177}
2178
2179pub(crate) fn zip_map4_into_validated<
2180 D: Copy + MaybeSendSync,
2181 A: Copy + MaybeSendSync,
2182 B: Copy + MaybeSendSync,
2183 C: Copy + MaybeSendSync,
2184 E: Copy + MaybeSendSync,
2185 OpA: ElementOp<A>,
2186 OpB: ElementOp<B>,
2187 OpC: ElementOp<C>,
2188 OpE: ElementOp<E>,
2189>(
2190 dest: &mut StridedViewMut<D>,
2191 a: &StridedView<A, OpA>,
2192 b: &StridedView<B, OpB>,
2193 c: &StridedView<C, OpC>,
2194 e: &StridedView<E, OpE>,
2195 f: impl Fn(A, B, C, E) -> D + MaybeSync,
2196 _validated: ValidatedDestinationLayout,
2197) -> Result<()> {
2198 ensure_same_shape(dest.dims(), a.dims())?;
2199 ensure_same_shape(dest.dims(), b.dims())?;
2200 ensure_same_shape(dest.dims(), c.dims())?;
2201 ensure_same_shape(dest.dims(), e.dims())?;
2202 let dst_ptr = dest.as_mut_ptr();
2203 let a_ptr = a.ptr();
2204 let b_ptr = b.ptr();
2205 let c_ptr = c.ptr();
2206 let e_ptr = e.ptr();
2207
2208 let dst_dims = dest.dims();
2209 let dst_strides = dest.strides();
2210
2211 if sequential_contiguous_layout(
2212 dst_dims,
2213 &[
2214 dst_strides,
2215 a.strides(),
2216 b.strides(),
2217 c.strides(),
2218 e.strides(),
2219 ],
2220 )?
2221 .is_some()
2222 {
2223 let len = total_len(dst_dims)?;
2224 let dst = unsafe { std::slice::from_raw_parts_mut(dst_ptr, len) };
2225 let sa = unsafe { std::slice::from_raw_parts(a_ptr, len) };
2226 let sb = unsafe { std::slice::from_raw_parts(b_ptr, len) };
2227 let sc = unsafe { std::slice::from_raw_parts(c_ptr, len) };
2228 let se = unsafe { std::slice::from_raw_parts(e_ptr, len) };
2229 simd::dispatch_if_large(len, || {
2230 for i in 0..len {
2231 dst[i] = f(
2232 OpA::apply(sa[i]),
2233 OpB::apply(sb[i]),
2234 OpC::apply(sc[i]),
2235 OpE::apply(se[i]),
2236 );
2237 }
2238 });
2239 return Ok(());
2240 }
2241
2242 let strides_list: [&[isize]; 5] = [
2243 dst_strides,
2244 a.strides(),
2245 b.strides(),
2246 c.strides(),
2247 e.strides(),
2248 ];
2249 let elem_size = std::mem::size_of::<D>()
2250 .max(std::mem::size_of::<A>())
2251 .max(std::mem::size_of::<B>())
2252 .max(std::mem::size_of::<C>())
2253 .max(std::mem::size_of::<E>());
2254 let total = total_len(dst_dims)?;
2255
2256 let (fused_dims, ordered_strides, plan) = if total <= SMALL_TENSOR_THRESHOLD {
2258 build_plan_fused_small(dst_dims, &strides_list)
2259 } else {
2260 build_plan_fused(dst_dims, &strides_list, Some(0), elem_size)
2261 };
2262
2263 #[cfg(feature = "parallel")]
2264 {
2265 let total = total_len(&fused_dims)?;
2266 let nthreads = crate::execution_policy::rayon_threads();
2267 if total > MINTHREADLENGTH && nthreads > 1 {
2268 use crate::threading::SendPtr;
2269 let dst_send = SendPtr(dst_ptr);
2270 let a_send = SendPtr(a_ptr as *mut A);
2271 let b_send = SendPtr(b_ptr as *mut B);
2272 let c_send = SendPtr(c_ptr as *mut C);
2273 let e_send = SendPtr(e_ptr as *mut E);
2274
2275 let costs = compute_costs(&ordered_strides);
2276 let initial_offsets = vec![0isize; strides_list.len()];
2277 return mapreduce_threaded(
2278 &fused_dims,
2279 &plan.block,
2280 &ordered_strides,
2281 &initial_offsets,
2282 &costs,
2283 nthreads,
2284 0,
2285 1,
2286 &|dims, blocks, strides_list, offsets| {
2287 for_each_inner_block_with_offsets(
2288 dims,
2289 blocks,
2290 strides_list,
2291 offsets,
2292 |offsets, len, strides| {
2293 let dp = unsafe { dst_send.as_ptr().offset(offsets[0]) };
2294 let ap = unsafe { a_send.as_const().offset(offsets[1]) };
2295 let bp = unsafe { b_send.as_const().offset(offsets[2]) };
2296 let cp = unsafe { c_send.as_const().offset(offsets[3]) };
2297 let ep = unsafe { e_send.as_const().offset(offsets[4]) };
2298 unsafe {
2299 inner_loop_map4::<D, A, B, C, E, OpA, OpB, OpC, OpE>(
2300 dp, strides[0], ap, strides[1], bp, strides[2], cp, strides[3],
2301 ep, strides[4], len, &f,
2302 )
2303 };
2304 Ok(())
2305 },
2306 )
2307 },
2308 );
2309 }
2310 }
2311
2312 let initial_offsets = vec![0isize; ordered_strides.len()];
2313 for_each_inner_block_preordered(
2314 &fused_dims,
2315 &plan.block,
2316 &ordered_strides,
2317 &initial_offsets,
2318 |offsets, len, strides| {
2319 let dp = unsafe { dst_ptr.offset(offsets[0]) };
2320 let ap = unsafe { a_ptr.offset(offsets[1]) };
2321 let bp = unsafe { b_ptr.offset(offsets[2]) };
2322 let cp = unsafe { c_ptr.offset(offsets[3]) };
2323 let ep = unsafe { e_ptr.offset(offsets[4]) };
2324 unsafe {
2325 inner_loop_map4::<D, A, B, C, E, OpA, OpB, OpC, OpE>(
2326 dp, strides[0], ap, strides[1], bp, strides[2], cp, strides[3], ep, strides[4],
2327 len, &f,
2328 )
2329 };
2330 Ok(())
2331 },
2332 )
2333}
2334
2335#[cfg(test)]
2336#[path = "map_view/tests/scalar_branch_tests.rs"]
2337mod scalar_branch_tests;
2338
2339#[cfg(all(test, feature = "parallel"))]
2340#[path = "map_view/tests/tests.rs"]
2341mod tests;