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 if self.coords[i] + 1 < dim {
662 self.coords[i] += 1;
663 self.a_offset += self.a_strides[axis];
664 self.b_offset += self.b_strides[axis];
665 break;
666 }
667
668 let last = (dim - 1) as isize;
669 self.coords[i] = 0;
670 self.a_offset -= last * self.a_strides[axis];
671 self.b_offset -= last * self.b_strides[axis];
672 }
673 }
674}
675
676#[inline(always)]
677unsafe fn run_contiguous_mul_row_block<
678 O: MulOutput<D>,
679 D: Copy + 'static,
680 A: Copy + Mul<B, Output = D> + 'static,
681 B: Copy + 'static,
682>(
683 dst_ptr: *mut O::Slot,
684 a_ptr: *const A,
685 b_ptr: *const B,
686 plan: &ContiguousMulRangePlan,
687 base_index: usize,
688 total: usize,
689 base_a_offset: isize,
690 base_b_offset: isize,
691) {
692 let inner_len = plan.inner_len.max(1);
693 let row_len = plan.row_len.max(1);
694 #[cfg(feature = "parallel")]
695 let block_len = inner_len.saturating_mul(row_len);
696
697 #[cfg(feature = "parallel")]
698 if total.saturating_sub(base_index) >= block_len {
699 match transposed_scalar_tile_kind(plan) {
700 Some(TransposedScalarTileKind::RhsScalar) => {
701 if simd::try_mul_transposed_scalar_rhs_2d::<D, A, B>(
702 dst_ptr.add(base_index).cast::<D>(),
703 a_ptr.offset(base_a_offset),
704 b_ptr.offset(base_b_offset),
705 inner_len,
706 row_len,
707 plan.a_fast_stride,
708 plan.a_row_stride,
709 ) {
710 return;
711 }
712 }
713 Some(TransposedScalarTileKind::LhsScalar) => {
714 if simd::try_mul_transposed_scalar_lhs_2d::<D, A, B>(
715 dst_ptr.add(base_index).cast::<D>(),
716 a_ptr.offset(base_a_offset),
717 b_ptr.offset(base_b_offset),
718 inner_len,
719 row_len,
720 plan.b_fast_stride,
721 plan.b_row_stride,
722 ) {
723 return;
724 }
725 }
726 None => {}
727 }
728 }
729
730 let mut index = base_index;
731 let mut a_offset = base_a_offset;
732 let mut b_offset = base_b_offset;
733
734 for _ in 0..row_len {
735 if index >= total {
736 break;
737 }
738 let len = inner_len.min(total - index);
739 inner_loop_mul2::<O, D, A, B>(
740 dst_ptr.add(index),
741 1,
742 a_ptr.offset(a_offset),
743 plan.a_fast_stride,
744 b_ptr.offset(b_offset),
745 plan.b_fast_stride,
746 len,
747 );
748 index += inner_len;
749 a_offset += plan.a_row_stride;
750 b_offset += plan.b_row_stride;
751 }
752}
753
754#[cfg(feature = "parallel")]
755fn strided_offset_for_contiguous_linear_index(
756 dims: &[usize],
757 strides: &[isize],
758 axis_order: &[usize],
759 mut index: usize,
760) -> isize {
761 let mut offset = 0isize;
762 for &axis in axis_order {
763 let dim = dims[axis];
764 if dim == 0 {
765 return 0;
766 }
767 let coord = index % dim;
768 index /= dim;
769 offset += coord as isize * strides[axis];
770 }
771 offset
772}
773
774fn try_contiguous_range_mul<
775 O: MulOutput<D>,
776 D: Copy + MaybeSendSync + 'static,
777 A: Copy + MaybeSendSync + Mul<B, Output = D> + 'static,
778 B: Copy + MaybeSendSync + 'static,
779>(
780 dst_ptr: *mut O::Slot,
781 dims: &[usize],
782 dst_strides: &[isize],
783 a_ptr: *const A,
784 a_strides: &[isize],
785 b_ptr: *const B,
786 b_strides: &[isize],
787) -> bool {
788 let Ok(total) = total_len(dims) else {
791 return false;
792 };
793 if total == 0 {
794 return true;
795 }
796 if total <= CONTIGUOUS_RANGE_MIN_LEN {
797 return false;
798 }
799
800 let Some(plan) = contiguous_mul_range_plan(dims, dst_strides, a_strides, b_strides) else {
801 return false;
802 };
803
804 let inner_len = plan.inner_len.max(1);
805 let row_len = plan.row_len.max(1);
806 let block_len = inner_len.saturating_mul(row_len).max(1);
807 let outer_groups = total.div_ceil(block_len);
808
809 #[cfg(feature = "parallel")]
810 {
811 let nthreads = crate::execution_policy::rayon_threads();
812 if nthreads > 1 {
813 use crate::threading::{parallel_for_each, SendPtr};
814
815 let dst = SendPtr(dst_ptr);
816 let a = SendPtr(a_ptr as *mut A);
817 let b = SendPtr(b_ptr as *mut B);
818
819 if outer_groups < nthreads {
820 let chunk_len = total.div_ceil(nthreads);
821 let nchunks = total.div_ceil(chunk_len);
822
823 parallel_for_each(0..nchunks, nthreads, &|chunks| {
824 for chunk in chunks {
825 let start = chunk * chunk_len;
826 let end = (start + chunk_len).min(total);
827 let mut index = start;
828
829 while index < end {
830 let in_inner = index % inner_len;
831 let len = (inner_len - in_inner).min(end - index);
832 let a_offset = strided_offset_for_contiguous_linear_index(
833 dims,
834 a_strides,
835 &plan.axis_order,
836 index,
837 );
838 let b_offset = strided_offset_for_contiguous_linear_index(
839 dims,
840 b_strides,
841 &plan.axis_order,
842 index,
843 );
844
845 unsafe {
846 inner_loop_mul2::<O, D, A, B>(
847 dst.as_ptr().add(index),
848 1,
849 a.as_const().offset(a_offset),
850 plan.a_fast_stride,
851 b.as_const().offset(b_offset),
852 plan.b_fast_stride,
853 len,
854 );
855 }
856 index += len;
857 }
858 }
859 });
860
861 return true;
862 }
863
864 let groups_per_chunk = outer_groups.div_ceil(nthreads);
865 let nchunks = outer_groups.div_ceil(groups_per_chunk);
866
867 parallel_for_each(0..nchunks, nthreads, &|chunks| {
868 for chunk in chunks {
869 let group_start = chunk * groups_per_chunk;
870 let group_end = (group_start + groups_per_chunk).min(outer_groups);
871 let mut cursor = ContiguousMulOuterCursor::new(
872 dims,
873 a_strides,
874 b_strides,
875 &plan,
876 group_start,
877 );
878
879 for group in group_start..group_end {
880 let index = group * block_len;
881 unsafe {
882 run_contiguous_mul_row_block::<O, D, A, B>(
883 dst.as_ptr(),
884 a.as_const(),
885 b.as_const(),
886 &plan,
887 index,
888 total,
889 cursor.a_offset,
890 cursor.b_offset,
891 );
892 }
893 cursor.advance();
894 }
895 }
896 });
897
898 true
899 } else {
900 run_contiguous_range_mul_single_thread::<O, D, A, B>(
901 dst_ptr,
902 dims,
903 a_ptr,
904 a_strides,
905 b_ptr,
906 b_strides,
907 &plan,
908 total,
909 block_len,
910 outer_groups,
911 )
912 }
913 }
914
915 #[cfg(not(feature = "parallel"))]
916 {
917 run_contiguous_range_mul_single_thread::<O, D, A, B>(
918 dst_ptr,
919 dims,
920 a_ptr,
921 a_strides,
922 b_ptr,
923 b_strides,
924 &plan,
925 total,
926 block_len,
927 outer_groups,
928 )
929 }
930}
931
932fn run_contiguous_range_mul_single_thread<
933 O: MulOutput<D>,
934 D: Copy + MaybeSendSync + 'static,
935 A: Copy + MaybeSendSync + Mul<B, Output = D> + 'static,
936 B: Copy + MaybeSendSync + 'static,
937>(
938 dst_ptr: *mut O::Slot,
939 dims: &[usize],
940 a_ptr: *const A,
941 a_strides: &[isize],
942 b_ptr: *const B,
943 b_strides: &[isize],
944 plan: &ContiguousMulRangePlan,
945 total: usize,
946 block_len: usize,
947 outer_groups: usize,
948) -> bool {
949 let mut cursor = ContiguousMulOuterCursor::new(dims, a_strides, b_strides, plan, 0);
950 for group in 0..outer_groups {
951 let index = group * block_len;
952 unsafe {
953 run_contiguous_mul_row_block::<O, D, A, B>(
954 dst_ptr,
955 a_ptr,
956 b_ptr,
957 plan,
958 index,
959 total,
960 cursor.a_offset,
961 cursor.b_offset,
962 );
963 }
964 cursor.advance();
965 }
966 true
967}
968
969#[inline(always)]
971unsafe fn inner_loop_map3<
972 D: Copy,
973 A: Copy,
974 B: Copy,
975 C: Copy,
976 OpA: ElementOp<A>,
977 OpB: ElementOp<B>,
978 OpC: ElementOp<C>,
979>(
980 dp: *mut D,
981 ds: isize,
982 ap: *const A,
983 a_s: isize,
984 bp: *const B,
985 b_s: isize,
986 cp: *const C,
987 c_s: isize,
988 len: usize,
989 f: &impl Fn(A, B, C) -> D,
990) {
991 if ds == 1 && a_s == 1 && b_s == 1 && c_s == 1 {
992 let src_a = std::slice::from_raw_parts(ap, len);
993 let src_b = std::slice::from_raw_parts(bp, len);
994 let src_c = std::slice::from_raw_parts(cp, len);
995 let dst = std::slice::from_raw_parts_mut(dp, len);
996 simd::dispatch_if_large(len, || {
997 for i in 0..len {
998 dst[i] = f(
999 OpA::apply(src_a[i]),
1000 OpB::apply(src_b[i]),
1001 OpC::apply(src_c[i]),
1002 );
1003 }
1004 });
1005 } else {
1006 let mut dp = dp;
1007 let mut ap = ap;
1008 let mut bp = bp;
1009 let mut cp = cp;
1010 for _ in 0..len {
1011 *dp = f(OpA::apply(*ap), OpB::apply(*bp), OpC::apply(*cp));
1012 dp = dp.offset(ds);
1013 ap = ap.offset(a_s);
1014 bp = bp.offset(b_s);
1015 cp = cp.offset(c_s);
1016 }
1017 }
1018}
1019
1020#[inline(always)]
1022unsafe fn inner_loop_map4<
1023 D: Copy,
1024 A: Copy,
1025 B: Copy,
1026 C: Copy,
1027 E: Copy,
1028 OpA: ElementOp<A>,
1029 OpB: ElementOp<B>,
1030 OpC: ElementOp<C>,
1031 OpE: ElementOp<E>,
1032>(
1033 dp: *mut D,
1034 ds: isize,
1035 ap: *const A,
1036 a_s: isize,
1037 bp: *const B,
1038 b_s: isize,
1039 cp: *const C,
1040 c_s: isize,
1041 ep: *const E,
1042 e_s: isize,
1043 len: usize,
1044 f: &impl Fn(A, B, C, E) -> D,
1045) {
1046 if ds == 1 && a_s == 1 && b_s == 1 && c_s == 1 && e_s == 1 {
1047 let src_a = std::slice::from_raw_parts(ap, len);
1048 let src_b = std::slice::from_raw_parts(bp, len);
1049 let src_c = std::slice::from_raw_parts(cp, len);
1050 let src_e = std::slice::from_raw_parts(ep, len);
1051 let dst = std::slice::from_raw_parts_mut(dp, len);
1052 simd::dispatch_if_large(len, || {
1053 for i in 0..len {
1054 dst[i] = f(
1055 OpA::apply(src_a[i]),
1056 OpB::apply(src_b[i]),
1057 OpC::apply(src_c[i]),
1058 OpE::apply(src_e[i]),
1059 );
1060 }
1061 });
1062 } else {
1063 let mut dp = dp;
1064 let mut ap = ap;
1065 let mut bp = bp;
1066 let mut cp = cp;
1067 let mut ep = ep;
1068 for _ in 0..len {
1069 *dp = f(
1070 OpA::apply(*ap),
1071 OpB::apply(*bp),
1072 OpC::apply(*cp),
1073 OpE::apply(*ep),
1074 );
1075 dp = dp.offset(ds);
1076 ap = ap.offset(a_s);
1077 bp = bp.offset(b_s);
1078 cp = cp.offset(c_s);
1079 ep = ep.offset(e_s);
1080 }
1081 }
1082}
1083
1084pub fn map_into<D: Copy + MaybeSendSync, A: Copy + MaybeSendSync, Op: ElementOp<A>>(
1089 dest: &mut StridedViewMut<D>,
1090 src: &StridedView<A, Op>,
1091 f: impl Fn(A) -> D + MaybeSync,
1092) -> Result<()> {
1093 map_parts_into::<D, A, Op>(
1094 dest.as_mut_ptr(),
1095 dest.dims(),
1096 dest.strides(),
1097 src.ptr(),
1098 src.dims(),
1099 src.strides(),
1100 f,
1101 )
1102}
1103
1104pub(crate) fn map_into_validated<
1105 D: Copy + MaybeSendSync,
1106 A: Copy + MaybeSendSync,
1107 Op: ElementOp<A>,
1108>(
1109 dest: &mut StridedViewMut<D>,
1110 src: &StridedView<A, Op>,
1111 f: impl Fn(A) -> D + MaybeSync,
1112 validated: ValidatedDestinationLayout,
1113) -> Result<()> {
1114 ensure_same_shape(dest.dims(), src.dims())?;
1115 map_parts_into_validated::<D, A, Op>(
1116 dest.as_mut_ptr(),
1117 dest.dims(),
1118 dest.strides(),
1119 src.ptr(),
1120 src.strides(),
1121 f,
1122 validated,
1123 )
1124}
1125
1126pub(crate) fn map_raw_into<D: Copy + MaybeSendSync, A: Copy + MaybeSendSync, Op: ElementOp<A>>(
1127 dest: &mut crate::RawStridedMut<'_, D>,
1128 src: &crate::RawStridedRef<'_, A>,
1129 f: impl Fn(A) -> D + MaybeSync,
1130) -> Result<()> {
1131 map_parts_into::<D, A, Op>(
1132 dest.as_mut_ptr(),
1133 dest.dims(),
1134 dest.strides(),
1135 src.ptr(),
1136 src.dims(),
1137 src.strides(),
1138 f,
1139 )
1140}
1141
1142pub(crate) fn map_raw_into_validated<
1143 D: Copy + MaybeSendSync,
1144 A: Copy + MaybeSendSync,
1145 Op: ElementOp<A>,
1146>(
1147 dest: &mut crate::RawStridedMut<'_, D>,
1148 src: &crate::RawStridedRef<'_, A>,
1149 f: impl Fn(A) -> D + MaybeSync,
1150 validated: ValidatedDestinationLayout,
1151) -> Result<()> {
1152 ensure_same_shape(dest.dims(), src.dims())?;
1153 map_parts_into_validated::<D, A, Op>(
1154 dest.as_mut_ptr(),
1155 dest.dims(),
1156 dest.strides(),
1157 src.ptr(),
1158 src.strides(),
1159 f,
1160 validated,
1161 )
1162}
1163
1164#[allow(clippy::too_many_arguments)]
1165fn map_parts_into<D: Copy + MaybeSendSync, A: Copy + MaybeSendSync, Op: ElementOp<A>>(
1166 dst_ptr: *mut D,
1167 dst_dims: &[usize],
1168 dst_strides: &[isize],
1169 src_ptr: *const A,
1170 src_dims: &[usize],
1171 src_strides: &[isize],
1172 f: impl Fn(A) -> D + MaybeSync,
1173) -> Result<()> {
1174 ensure_same_shape(dst_dims, src_dims)?;
1175 let validated = validate_destination_layout(dst_dims, dst_strides)?;
1176 map_parts_into_validated::<D, A, Op>(
1177 dst_ptr,
1178 dst_dims,
1179 dst_strides,
1180 src_ptr,
1181 src_strides,
1182 f,
1183 validated,
1184 )
1185}
1186
1187#[allow(clippy::too_many_arguments)]
1188fn map_parts_into_validated<D: Copy + MaybeSendSync, A: Copy + MaybeSendSync, Op: ElementOp<A>>(
1189 dst_ptr: *mut D,
1190 dst_dims: &[usize],
1191 dst_strides: &[isize],
1192 src_ptr: *const A,
1193 src_strides: &[isize],
1194 f: impl Fn(A) -> D + MaybeSync,
1195 _validated: ValidatedDestinationLayout,
1196) -> Result<()> {
1197 if sequential_contiguous_layout(dst_dims, &[dst_strides, src_strides])?.is_some() {
1198 let len = total_len(dst_dims)?;
1199 let dst = unsafe { std::slice::from_raw_parts_mut(dst_ptr, len) };
1200 let src = unsafe { std::slice::from_raw_parts(src_ptr, len) };
1201 simd::dispatch_if_large(len, || {
1202 for i in 0..len {
1203 dst[i] = f(Op::apply(src[i]));
1204 }
1205 });
1206 return Ok(());
1207 }
1208
1209 let strides_list: [&[isize]; 2] = [dst_strides, src_strides];
1210 let elem_size = std::mem::size_of::<D>().max(std::mem::size_of::<A>());
1211 let total = total_len(dst_dims)?;
1212
1213 let (fused_dims, ordered_strides, plan) = if total <= SMALL_TENSOR_THRESHOLD {
1215 build_plan_fused_small(dst_dims, &strides_list)
1216 } else {
1217 build_plan_fused(dst_dims, &strides_list, Some(0), elem_size)
1218 };
1219
1220 #[cfg(feature = "parallel")]
1221 {
1222 let total = total_len(&fused_dims)?;
1223 let nthreads = crate::execution_policy::rayon_threads();
1224 if total > MINTHREADLENGTH && nthreads > 1 {
1225 use crate::threading::SendPtr;
1226 let dst_send = SendPtr(dst_ptr);
1227 let src_send = SendPtr(src_ptr as *mut A);
1228
1229 let costs = compute_costs(&ordered_strides);
1230 let initial_offsets = vec![0isize; strides_list.len()];
1231 return mapreduce_threaded(
1232 &fused_dims,
1233 &plan.block,
1234 &ordered_strides,
1235 &initial_offsets,
1236 &costs,
1237 nthreads,
1238 0,
1239 1,
1240 &|dims, blocks, strides_list, offsets| {
1241 for_each_inner_block_with_offsets(
1242 dims,
1243 blocks,
1244 strides_list,
1245 offsets,
1246 |offsets, len, strides| {
1247 let dp = unsafe { dst_send.as_ptr().offset(offsets[0]) };
1248 let sp = unsafe { src_send.as_const().offset(offsets[1]) };
1249 unsafe {
1250 inner_loop_map1::<D, A, Op>(dp, strides[0], sp, strides[1], len, &f)
1251 };
1252 Ok(())
1253 },
1254 )
1255 },
1256 );
1257 }
1258 }
1259
1260 let initial_offsets = vec![0isize; ordered_strides.len()];
1261 for_each_inner_block_preordered(
1262 &fused_dims,
1263 &plan.block,
1264 &ordered_strides,
1265 &initial_offsets,
1266 |offsets, len, strides| {
1267 let dp = unsafe { dst_ptr.offset(offsets[0]) };
1268 let sp = unsafe { src_ptr.offset(offsets[1]) };
1269 unsafe { inner_loop_map1::<D, A, Op>(dp, strides[0], sp, strides[1], len, &f) };
1270 Ok(())
1271 },
1272 )
1273}
1274
1275pub fn zip_map2_into<
1280 D: Copy + MaybeSendSync,
1281 A: Copy + MaybeSendSync,
1282 B: Copy + MaybeSendSync,
1283 OpA: ElementOp<A>,
1284 OpB: ElementOp<B>,
1285>(
1286 dest: &mut StridedViewMut<D>,
1287 a: &StridedView<A, OpA>,
1288 b: &StridedView<B, OpB>,
1289 f: impl Fn(A, B) -> D + MaybeSync,
1290) -> Result<()> {
1291 zip_map2_parts_into::<D, A, B, OpA, OpB>(
1292 dest.as_mut_ptr(),
1293 dest.dims(),
1294 dest.strides(),
1295 a.ptr(),
1296 a.dims(),
1297 a.strides(),
1298 b.ptr(),
1299 b.dims(),
1300 b.strides(),
1301 f,
1302 )
1303}
1304
1305pub(crate) fn zip_map2_into_validated<
1306 D: Copy + MaybeSendSync,
1307 A: Copy + MaybeSendSync,
1308 B: Copy + MaybeSendSync,
1309 OpA: ElementOp<A>,
1310 OpB: ElementOp<B>,
1311>(
1312 dest: &mut StridedViewMut<D>,
1313 a: &StridedView<A, OpA>,
1314 b: &StridedView<B, OpB>,
1315 f: impl Fn(A, B) -> D + MaybeSync,
1316 validated: ValidatedDestinationLayout,
1317) -> Result<()> {
1318 zip_map2_parts_into_validated::<D, A, B, OpA, OpB>(
1319 dest.as_mut_ptr(),
1320 dest.dims(),
1321 dest.strides(),
1322 a.ptr(),
1323 a.strides(),
1324 b.ptr(),
1325 b.strides(),
1326 f,
1327 validated,
1328 )
1329}
1330
1331#[non_exhaustive]
1333#[derive(Clone, Copy, Debug, Eq, PartialEq)]
1334pub enum CompareOp {
1335 Eq,
1336 Lt,
1337 Le,
1338 Gt,
1339 Ge,
1340}
1341
1342pub fn compare_into<T, OpA, OpB>(
1353 dest: &mut StridedViewMut<bool>,
1354 a: &StridedView<T, OpA>,
1355 b: &StridedView<T, OpB>,
1356 op: CompareOp,
1357) -> Result<()>
1358where
1359 T: Copy + MaybeSendSync + PartialOrd,
1360 OpA: ElementOp<T>,
1361 OpB: ElementOp<T>,
1362{
1363 match op {
1364 CompareOp::Eq => zip_map2_into(dest, a, b, |lhs, rhs| lhs == rhs),
1365 CompareOp::Lt => zip_map2_into(dest, a, b, |lhs, rhs| lhs < rhs),
1366 CompareOp::Le => zip_map2_into(dest, a, b, |lhs, rhs| lhs <= rhs),
1367 CompareOp::Gt => zip_map2_into(dest, a, b, |lhs, rhs| lhs > rhs),
1368 CompareOp::Ge => zip_map2_into(dest, a, b, |lhs, rhs| lhs >= rhs),
1369 }
1370}
1371
1372pub fn compare_into_uninit<T, OpA, OpB>(
1391 dest: &mut StridedViewMut<MaybeUninit<bool>>,
1392 a: &StridedView<T, OpA>,
1393 b: &StridedView<T, OpB>,
1394 op: CompareOp,
1395) -> Result<()>
1396where
1397 T: Copy + MaybeSendSync + PartialOrd,
1398 OpA: ElementOp<T>,
1399 OpB: ElementOp<T>,
1400{
1401 ensure_same_shape(dest.dims(), a.dims())?;
1402 ensure_same_shape(dest.dims(), b.dims())?;
1403 let validated = validate_destination_layout(dest.dims(), dest.strides())?;
1404 validate_typed_no_overlap(dest, a, 0)?;
1405 validate_typed_no_overlap(dest, b, 1)?;
1406 match op {
1407 CompareOp::Eq => zip_map2_into_validated(
1408 dest,
1409 a,
1410 b,
1411 |lhs, rhs| MaybeUninit::new(lhs == rhs),
1412 validated,
1413 ),
1414 CompareOp::Lt => zip_map2_into_validated(
1415 dest,
1416 a,
1417 b,
1418 |lhs, rhs| MaybeUninit::new(lhs < rhs),
1419 validated,
1420 ),
1421 CompareOp::Le => zip_map2_into_validated(
1422 dest,
1423 a,
1424 b,
1425 |lhs, rhs| MaybeUninit::new(lhs <= rhs),
1426 validated,
1427 ),
1428 CompareOp::Gt => zip_map2_into_validated(
1429 dest,
1430 a,
1431 b,
1432 |lhs, rhs| MaybeUninit::new(lhs > rhs),
1433 validated,
1434 ),
1435 CompareOp::Ge => zip_map2_into_validated(
1436 dest,
1437 a,
1438 b,
1439 |lhs, rhs| MaybeUninit::new(lhs >= rhs),
1440 validated,
1441 ),
1442 }
1443}
1444
1445pub(crate) fn zip_map2_raw_into_validated<
1446 D: Copy + MaybeSendSync,
1447 A: Copy + MaybeSendSync,
1448 B: Copy + MaybeSendSync,
1449 OpA: ElementOp<A>,
1450 OpB: ElementOp<B>,
1451>(
1452 dest: &mut crate::RawStridedMut<'_, D>,
1453 a: &crate::RawStridedRef<'_, A>,
1454 b: &crate::RawStridedRef<'_, B>,
1455 f: impl Fn(A, B) -> D + MaybeSync,
1456 validated: ValidatedDestinationLayout,
1457) -> Result<()> {
1458 ensure_same_shape(dest.dims(), a.dims())?;
1459 ensure_same_shape(dest.dims(), b.dims())?;
1460 zip_map2_parts_into_validated::<D, A, B, OpA, OpB>(
1461 dest.as_mut_ptr(),
1462 dest.dims(),
1463 dest.strides(),
1464 a.ptr(),
1465 a.strides(),
1466 b.ptr(),
1467 b.strides(),
1468 f,
1469 validated,
1470 )
1471}
1472
1473#[allow(clippy::too_many_arguments)]
1474fn zip_map2_parts_into<
1475 D: Copy + MaybeSendSync,
1476 A: Copy + MaybeSendSync,
1477 B: Copy + MaybeSendSync,
1478 OpA: ElementOp<A>,
1479 OpB: ElementOp<B>,
1480>(
1481 dst_ptr: *mut D,
1482 dst_dims: &[usize],
1483 dst_strides: &[isize],
1484 a_ptr: *const A,
1485 a_dims: &[usize],
1486 a_strides: &[isize],
1487 b_ptr: *const B,
1488 b_dims: &[usize],
1489 b_strides: &[isize],
1490 f: impl Fn(A, B) -> D + MaybeSync,
1491) -> Result<()> {
1492 ensure_same_shape(dst_dims, a_dims)?;
1493 ensure_same_shape(dst_dims, b_dims)?;
1494 let validated = validate_destination_layout(dst_dims, dst_strides)?;
1495 zip_map2_parts_into_validated::<D, A, B, OpA, OpB>(
1496 dst_ptr,
1497 dst_dims,
1498 dst_strides,
1499 a_ptr,
1500 a_strides,
1501 b_ptr,
1502 b_strides,
1503 f,
1504 validated,
1505 )
1506}
1507
1508#[allow(clippy::too_many_arguments)]
1509fn zip_map2_parts_into_validated<
1510 D: Copy + MaybeSendSync,
1511 A: Copy + MaybeSendSync,
1512 B: Copy + MaybeSendSync,
1513 OpA: ElementOp<A>,
1514 OpB: ElementOp<B>,
1515>(
1516 dst_ptr: *mut D,
1517 dst_dims: &[usize],
1518 dst_strides: &[isize],
1519 a_ptr: *const A,
1520 a_strides: &[isize],
1521 b_ptr: *const B,
1522 b_strides: &[isize],
1523 f: impl Fn(A, B) -> D + MaybeSync,
1524 _validated: ValidatedDestinationLayout,
1525) -> Result<()> {
1526 if sequential_contiguous_layout(dst_dims, &[dst_strides, a_strides, b_strides])?.is_some() {
1527 let len = total_len(dst_dims)?;
1528 let dst = unsafe { std::slice::from_raw_parts_mut(dst_ptr, len) };
1529 let sa = unsafe { std::slice::from_raw_parts(a_ptr, len) };
1530 let sb = unsafe { std::slice::from_raw_parts(b_ptr, len) };
1531 simd::dispatch_if_large(len, || {
1532 for i in 0..len {
1533 dst[i] = f(OpA::apply(sa[i]), OpB::apply(sb[i]));
1534 }
1535 });
1536 return Ok(());
1537 }
1538
1539 let strides_list: [&[isize]; 3] = [dst_strides, a_strides, b_strides];
1540 let elem_size = std::mem::size_of::<D>()
1541 .max(std::mem::size_of::<A>())
1542 .max(std::mem::size_of::<B>());
1543 let total = total_len(dst_dims)?;
1544
1545 let (fused_dims, ordered_strides, plan) = if total <= SMALL_TENSOR_THRESHOLD {
1547 build_plan_fused_small(dst_dims, &strides_list)
1548 } else {
1549 build_plan_fused(dst_dims, &strides_list, Some(0), elem_size)
1550 };
1551
1552 #[cfg(feature = "parallel")]
1553 {
1554 let total = total_len(&fused_dims)?;
1555 let nthreads = crate::execution_policy::rayon_threads();
1556 if total > MINTHREADLENGTH && nthreads > 1 {
1557 use crate::threading::SendPtr;
1558 let dst_send = SendPtr(dst_ptr);
1559 let a_send = SendPtr(a_ptr as *mut A);
1560 let b_send = SendPtr(b_ptr as *mut B);
1561
1562 let costs = compute_costs(&ordered_strides);
1563 let initial_offsets = vec![0isize; strides_list.len()];
1564 return mapreduce_threaded(
1565 &fused_dims,
1566 &plan.block,
1567 &ordered_strides,
1568 &initial_offsets,
1569 &costs,
1570 nthreads,
1571 0,
1572 1,
1573 &|dims, blocks, strides_list, offsets| {
1574 for_each_inner_block_with_offsets(
1575 dims,
1576 blocks,
1577 strides_list,
1578 offsets,
1579 |offsets, len, strides| {
1580 let dp = unsafe { dst_send.as_ptr().offset(offsets[0]) };
1581 let ap = unsafe { a_send.as_const().offset(offsets[1]) };
1582 let bp = unsafe { b_send.as_const().offset(offsets[2]) };
1583 unsafe {
1584 inner_loop_map2::<D, A, B, OpA, OpB>(
1585 dp, strides[0], ap, strides[1], bp, strides[2], len, &f,
1586 )
1587 };
1588 Ok(())
1589 },
1590 )
1591 },
1592 );
1593 }
1594 }
1595
1596 let initial_offsets = vec![0isize; ordered_strides.len()];
1597 for_each_inner_block_preordered(
1598 &fused_dims,
1599 &plan.block,
1600 &ordered_strides,
1601 &initial_offsets,
1602 |offsets, len, strides| {
1603 let dp = unsafe { dst_ptr.offset(offsets[0]) };
1604 let ap = unsafe { a_ptr.offset(offsets[1]) };
1605 let bp = unsafe { b_ptr.offset(offsets[2]) };
1606 unsafe {
1607 inner_loop_map2::<D, A, B, OpA, OpB>(
1608 dp, strides[0], ap, strides[1], bp, strides[2], len, &f,
1609 )
1610 };
1611 Ok(())
1612 },
1613 )
1614}
1615
1616fn mul_identity_into_raw<
1617 O: MulOutput<D>,
1618 D: Copy + MaybeSendSync + 'static,
1619 A: Copy + MaybeSendSync + Mul<B, Output = D> + 'static,
1620 B: Copy + MaybeSendSync + 'static,
1621>(
1622 dst_ptr: *mut O::Slot,
1623 dst_dims: &[usize],
1624 dst_strides: &[isize],
1625 a_ptr: *const A,
1626 a_strides: &[isize],
1627 b_ptr: *const B,
1628 b_strides: &[isize],
1629 _validated: ValidatedDestinationLayout,
1630) -> Result<()> {
1631 debug_assert_eq!(dst_dims.len(), a_strides.len());
1632 debug_assert_eq!(dst_dims.len(), b_strides.len());
1633
1634 if sequential_contiguous_layout(dst_dims, &[dst_strides, a_strides, b_strides])?.is_some() {
1635 let len = total_len(dst_dims)?;
1636 let sa = unsafe { std::slice::from_raw_parts(a_ptr, len) };
1637 let sb = unsafe { std::slice::from_raw_parts(b_ptr, len) };
1638 if unsafe { O::try_contiguous(dst_ptr, len, sa, sb) } {
1639 return Ok(());
1640 }
1641 for i in 0..len {
1642 unsafe { O::write(dst_ptr.add(i), multiply_value(sa[i], sb[i])) };
1643 }
1644 return Ok(());
1645 }
1646
1647 let strides_list: [&[isize]; 3] = [dst_strides, a_strides, b_strides];
1648 let elem_size = std::mem::size_of::<D>()
1649 .max(std::mem::size_of::<A>())
1650 .max(std::mem::size_of::<B>());
1651 let total = total_len(dst_dims)?;
1652
1653 if try_contiguous_range_mul::<O, D, A, B>(
1654 dst_ptr,
1655 dst_dims,
1656 dst_strides,
1657 a_ptr,
1658 a_strides,
1659 b_ptr,
1660 b_strides,
1661 ) {
1662 return Ok(());
1663 }
1664
1665 let (fused_dims, ordered_strides, plan) = if total <= SMALL_TENSOR_THRESHOLD {
1666 build_plan_fused_small(dst_dims, &strides_list)
1667 } else {
1668 build_plan_fused(dst_dims, &strides_list, Some(0), elem_size)
1669 };
1670
1671 #[cfg(feature = "parallel")]
1672 {
1673 let total = total_len(&fused_dims)?;
1674 let nthreads = crate::execution_policy::rayon_threads();
1675 if total > MINTHREADLENGTH && nthreads > 1 {
1676 use crate::threading::SendPtr;
1677 let dst_send = SendPtr(dst_ptr);
1678 let a_send = SendPtr(a_ptr as *mut A);
1679 let b_send = SendPtr(b_ptr as *mut B);
1680
1681 let costs = compute_costs(&ordered_strides);
1682 let initial_offsets = vec![0isize; strides_list.len()];
1683 return mapreduce_threaded(
1684 &fused_dims,
1685 &plan.block,
1686 &ordered_strides,
1687 &initial_offsets,
1688 &costs,
1689 nthreads,
1690 0,
1691 1,
1692 &|dims, blocks, strides_list, offsets| {
1693 for_each_inner_block_with_offsets(
1694 dims,
1695 blocks,
1696 strides_list,
1697 offsets,
1698 |offsets, len, strides| {
1699 let dp = unsafe { dst_send.as_ptr().offset(offsets[0]) };
1700 let ap = unsafe { a_send.as_const().offset(offsets[1]) };
1701 let bp = unsafe { b_send.as_const().offset(offsets[2]) };
1702 unsafe {
1703 inner_loop_mul2::<O, D, A, B>(
1704 dp, strides[0], ap, strides[1], bp, strides[2], len,
1705 )
1706 };
1707 Ok(())
1708 },
1709 )
1710 },
1711 );
1712 }
1713 }
1714
1715 let initial_offsets = vec![0isize; ordered_strides.len()];
1716 for_each_inner_block_preordered(
1717 &fused_dims,
1718 &plan.block,
1719 &ordered_strides,
1720 &initial_offsets,
1721 |offsets, len, strides| {
1722 let dp = unsafe { dst_ptr.offset(offsets[0]) };
1723 let ap = unsafe { a_ptr.offset(offsets[1]) };
1724 let bp = unsafe { b_ptr.offset(offsets[2]) };
1725 unsafe {
1726 inner_loop_mul2::<O, D, A, B>(dp, strides[0], ap, strides[1], bp, strides[2], len)
1727 };
1728 Ok(())
1729 },
1730 )
1731}
1732
1733pub fn mul_into<
1738 D: Copy + MaybeSendSync + 'static,
1739 A: Copy + Mul<B, Output = D> + MaybeSendSync + 'static,
1740 B: Copy + MaybeSendSync + 'static,
1741 OpA: ElementOp<A>,
1742 OpB: ElementOp<B>,
1743>(
1744 dest: &mut StridedViewMut<D>,
1745 a: &StridedView<A, OpA>,
1746 b: &StridedView<B, OpB>,
1747) -> Result<()> {
1748 ensure_same_shape(dest.dims(), a.dims())?;
1749 ensure_same_shape(dest.dims(), b.dims())?;
1750 let validated = validate_destination_layout(dest.dims(), dest.strides())?;
1751
1752 if OpA::IS_IDENTITY && OpB::IS_IDENTITY {
1753 return mul_identity_into_raw::<InitializedOutput, D, A, B>(
1754 dest.as_mut_ptr(),
1755 dest.dims(),
1756 dest.strides(),
1757 a.ptr(),
1758 a.strides(),
1759 b.ptr(),
1760 b.strides(),
1761 validated,
1762 );
1763 }
1764
1765 zip_map2_into_validated(dest, a, b, multiply_value, validated)
1766}
1767
1768pub fn mul_into_uninit<
1787 D: Copy + MaybeSendSync + 'static,
1788 A: Copy + Mul<B, Output = D> + MaybeSendSync + 'static,
1789 B: Copy + MaybeSendSync + 'static,
1790 OpA: ElementOp<A>,
1791 OpB: ElementOp<B>,
1792>(
1793 dest: &mut StridedViewMut<MaybeUninit<D>>,
1794 a: &StridedView<A, OpA>,
1795 b: &StridedView<B, OpB>,
1796) -> Result<()> {
1797 ensure_same_shape(dest.dims(), a.dims())?;
1798 ensure_same_shape(dest.dims(), b.dims())?;
1799 let validated = validate_destination_layout(dest.dims(), dest.strides())?;
1800 validate_typed_no_overlap(dest, a, 0)?;
1801 validate_typed_no_overlap(dest, b, 1)?;
1802
1803 if OpA::IS_IDENTITY && OpB::IS_IDENTITY {
1804 return mul_identity_into_raw::<UninitializedOutput, D, A, B>(
1805 dest.as_mut_ptr(),
1806 dest.dims(),
1807 dest.strides(),
1808 a.ptr(),
1809 a.strides(),
1810 b.ptr(),
1811 b.strides(),
1812 validated,
1813 );
1814 }
1815
1816 zip_map2_into_validated(
1817 dest,
1818 a,
1819 b,
1820 |lhs, rhs| MaybeUninit::new(multiply_value(lhs, rhs)),
1821 validated,
1822 )
1823}
1824
1825fn broadcast_strides_for_axes(
1826 source_dims: &[usize],
1827 source_strides: &[isize],
1828 target_dims: &[usize],
1829 axes: &[usize],
1830) -> Result<AxisVec<isize>> {
1831 if source_dims.len() != axes.len() {
1832 return Err(StridedError::RankMismatch(source_dims.len(), axes.len()));
1833 }
1834 debug_assert_eq!(source_dims.len(), source_strides.len());
1835
1836 let mut seen = AxisVec::<bool>::new();
1837 seen.resize(target_dims.len(), false);
1838 let mut strides = AxisVec::<isize>::new();
1839 strides.resize(target_dims.len(), 0);
1840 for (src_axis, &dst_axis) in axes.iter().enumerate() {
1841 if dst_axis >= target_dims.len() {
1842 return Err(StridedError::InvalidAxis {
1843 axis: dst_axis,
1844 rank: target_dims.len(),
1845 });
1846 }
1847 if seen[dst_axis] {
1848 return Err(StridedError::InvalidAxis {
1849 axis: dst_axis,
1850 rank: target_dims.len(),
1851 });
1852 }
1853 seen[dst_axis] = true;
1854
1855 let source_dim = source_dims[src_axis];
1856 let target_dim = target_dims[dst_axis];
1857 if source_dim != target_dim && source_dim != 1 {
1858 return Err(StridedError::ShapeMismatch(
1859 source_dims.to_vec(),
1860 target_dims.to_vec(),
1861 ));
1862 }
1863 if source_dim == target_dim {
1864 strides[dst_axis] = source_strides[src_axis];
1865 }
1866 }
1867
1868 Ok(strides)
1869}
1870
1871fn broadcast_view_with_strides<'a, T, Op: ElementOp<T>>(
1872 view: &StridedView<'a, T, Op>,
1873 target_dims: &[usize],
1874 strides: &[isize],
1875) -> StridedView<'a, T, Op> {
1876 unsafe { StridedView::new_unchecked(view.data(), target_dims, strides, view.offset()) }
1877}
1878
1879pub fn broadcast_mul_into<
1884 D: Copy + MaybeSendSync + 'static,
1885 A: Copy + Mul<B, Output = D> + MaybeSendSync + 'static,
1886 B: Copy + MaybeSendSync + 'static,
1887 OpA: ElementOp<A>,
1888 OpB: ElementOp<B>,
1889>(
1890 dest: &mut StridedViewMut<D>,
1891 a: &StridedView<A, OpA>,
1892 a_axes: &[usize],
1893 b: &StridedView<B, OpB>,
1894 b_axes: &[usize],
1895) -> Result<()> {
1896 let validated = validate_destination_layout(dest.dims(), dest.strides())?;
1897 let a_strides = broadcast_strides_for_axes(a.dims(), a.strides(), dest.dims(), a_axes)?;
1898 let b_strides = broadcast_strides_for_axes(b.dims(), b.strides(), dest.dims(), b_axes)?;
1899
1900 if OpA::IS_IDENTITY && OpB::IS_IDENTITY {
1901 return mul_identity_into_raw::<InitializedOutput, D, A, B>(
1902 dest.as_mut_ptr(),
1903 dest.dims(),
1904 dest.strides(),
1905 a.ptr(),
1906 &a_strides,
1907 b.ptr(),
1908 &b_strides,
1909 validated,
1910 );
1911 }
1912
1913 let a = broadcast_view_with_strides(a, dest.dims(), &a_strides);
1914 let b = broadcast_view_with_strides(b, dest.dims(), &b_strides);
1915 zip_map2_into_validated(dest, &a, &b, multiply_value, validated)
1916}
1917
1918pub fn broadcast_mul_into_uninit<
1937 D: Copy + MaybeSendSync + 'static,
1938 A: Copy + Mul<B, Output = D> + MaybeSendSync + 'static,
1939 B: Copy + MaybeSendSync + 'static,
1940 OpA: ElementOp<A>,
1941 OpB: ElementOp<B>,
1942>(
1943 dest: &mut StridedViewMut<MaybeUninit<D>>,
1944 a: &StridedView<A, OpA>,
1945 a_axes: &[usize],
1946 b: &StridedView<B, OpB>,
1947 b_axes: &[usize],
1948) -> Result<()> {
1949 let validated = validate_destination_layout(dest.dims(), dest.strides())?;
1950 let a_strides = broadcast_strides_for_axes(a.dims(), a.strides(), dest.dims(), a_axes)?;
1951 let b_strides = broadcast_strides_for_axes(b.dims(), b.strides(), dest.dims(), b_axes)?;
1952 validate_typed_no_overlap(dest, a, 0)?;
1953 validate_typed_no_overlap(dest, b, 1)?;
1954
1955 if OpA::IS_IDENTITY && OpB::IS_IDENTITY {
1956 return mul_identity_into_raw::<UninitializedOutput, D, A, B>(
1957 dest.as_mut_ptr(),
1958 dest.dims(),
1959 dest.strides(),
1960 a.ptr(),
1961 &a_strides,
1962 b.ptr(),
1963 &b_strides,
1964 validated,
1965 );
1966 }
1967
1968 let a = broadcast_view_with_strides(a, dest.dims(), &a_strides);
1969 let b = broadcast_view_with_strides(b, dest.dims(), &b_strides);
1970 zip_map2_into_validated(
1971 dest,
1972 &a,
1973 &b,
1974 |lhs, rhs| MaybeUninit::new(multiply_value(lhs, rhs)),
1975 validated,
1976 )
1977}
1978
1979pub fn zip_map3_into<
1981 D: Copy + MaybeSendSync,
1982 A: Copy + MaybeSendSync,
1983 B: Copy + MaybeSendSync,
1984 C: Copy + MaybeSendSync,
1985 OpA: ElementOp<A>,
1986 OpB: ElementOp<B>,
1987 OpC: ElementOp<C>,
1988>(
1989 dest: &mut StridedViewMut<D>,
1990 a: &StridedView<A, OpA>,
1991 b: &StridedView<B, OpB>,
1992 c: &StridedView<C, OpC>,
1993 f: impl Fn(A, B, C) -> D + MaybeSync,
1994) -> Result<()> {
1995 ensure_same_shape(dest.dims(), a.dims())?;
1996 ensure_same_shape(dest.dims(), b.dims())?;
1997 ensure_same_shape(dest.dims(), c.dims())?;
1998 let validated = validate_destination_layout(dest.dims(), dest.strides())?;
1999 zip_map3_into_validated(dest, a, b, c, f, validated)
2000}
2001
2002pub(crate) fn zip_map3_into_validated<
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 _validated: ValidatedDestinationLayout,
2017) -> Result<()> {
2018 ensure_same_shape(dest.dims(), a.dims())?;
2019 ensure_same_shape(dest.dims(), b.dims())?;
2020 ensure_same_shape(dest.dims(), c.dims())?;
2021 let dst_ptr = dest.as_mut_ptr();
2022 let a_ptr = a.ptr();
2023 let b_ptr = b.ptr();
2024 let c_ptr = c.ptr();
2025
2026 let dst_dims = dest.dims();
2027 let dst_strides = dest.strides();
2028
2029 if sequential_contiguous_layout(
2030 dst_dims,
2031 &[dst_strides, a.strides(), b.strides(), c.strides()],
2032 )?
2033 .is_some()
2034 {
2035 let len = total_len(dst_dims)?;
2036 let dst = unsafe { std::slice::from_raw_parts_mut(dst_ptr, len) };
2037 let sa = unsafe { std::slice::from_raw_parts(a_ptr, len) };
2038 let sb = unsafe { std::slice::from_raw_parts(b_ptr, len) };
2039 let sc = unsafe { std::slice::from_raw_parts(c_ptr, len) };
2040 simd::dispatch_if_large(len, || {
2041 for i in 0..len {
2042 dst[i] = f(OpA::apply(sa[i]), OpB::apply(sb[i]), OpC::apply(sc[i]));
2043 }
2044 });
2045 return Ok(());
2046 }
2047
2048 let strides_list: [&[isize]; 4] = [dst_strides, a.strides(), b.strides(), c.strides()];
2049 let elem_size = std::mem::size_of::<D>()
2050 .max(std::mem::size_of::<A>())
2051 .max(std::mem::size_of::<B>())
2052 .max(std::mem::size_of::<C>());
2053 let total = total_len(dst_dims)?;
2054
2055 let (fused_dims, ordered_strides, plan) = if total <= SMALL_TENSOR_THRESHOLD {
2057 build_plan_fused_small(dst_dims, &strides_list)
2058 } else {
2059 build_plan_fused(dst_dims, &strides_list, Some(0), elem_size)
2060 };
2061
2062 #[cfg(feature = "parallel")]
2063 {
2064 let total = total_len(&fused_dims)?;
2065 let nthreads = crate::execution_policy::rayon_threads();
2066 if total > MINTHREADLENGTH && nthreads > 1 {
2067 use crate::threading::SendPtr;
2068 let dst_send = SendPtr(dst_ptr);
2069 let a_send = SendPtr(a_ptr as *mut A);
2070 let b_send = SendPtr(b_ptr as *mut B);
2071 let c_send = SendPtr(c_ptr as *mut C);
2072
2073 let costs = compute_costs(&ordered_strides);
2074 let initial_offsets = vec![0isize; strides_list.len()];
2075 return mapreduce_threaded(
2076 &fused_dims,
2077 &plan.block,
2078 &ordered_strides,
2079 &initial_offsets,
2080 &costs,
2081 nthreads,
2082 0,
2083 1,
2084 &|dims, blocks, strides_list, offsets| {
2085 for_each_inner_block_with_offsets(
2086 dims,
2087 blocks,
2088 strides_list,
2089 offsets,
2090 |offsets, len, strides| {
2091 let dp = unsafe { dst_send.as_ptr().offset(offsets[0]) };
2092 let ap = unsafe { a_send.as_const().offset(offsets[1]) };
2093 let bp = unsafe { b_send.as_const().offset(offsets[2]) };
2094 let cp = unsafe { c_send.as_const().offset(offsets[3]) };
2095 unsafe {
2096 inner_loop_map3::<D, A, B, C, OpA, OpB, OpC>(
2097 dp, strides[0], ap, strides[1], bp, strides[2], cp, strides[3],
2098 len, &f,
2099 )
2100 };
2101 Ok(())
2102 },
2103 )
2104 },
2105 );
2106 }
2107 }
2108
2109 let initial_offsets = vec![0isize; ordered_strides.len()];
2110 for_each_inner_block_preordered(
2111 &fused_dims,
2112 &plan.block,
2113 &ordered_strides,
2114 &initial_offsets,
2115 |offsets, len, strides| {
2116 let dp = unsafe { dst_ptr.offset(offsets[0]) };
2117 let ap = unsafe { a_ptr.offset(offsets[1]) };
2118 let bp = unsafe { b_ptr.offset(offsets[2]) };
2119 let cp = unsafe { c_ptr.offset(offsets[3]) };
2120 unsafe {
2121 inner_loop_map3::<D, A, B, C, OpA, OpB, OpC>(
2122 dp, strides[0], ap, strides[1], bp, strides[2], cp, strides[3], len, &f,
2123 )
2124 };
2125 Ok(())
2126 },
2127 )
2128}
2129
2130pub fn zip_map4_into<
2132 D: Copy + MaybeSendSync,
2133 A: Copy + MaybeSendSync,
2134 B: Copy + MaybeSendSync,
2135 C: Copy + MaybeSendSync,
2136 E: Copy + MaybeSendSync,
2137 OpA: ElementOp<A>,
2138 OpB: ElementOp<B>,
2139 OpC: ElementOp<C>,
2140 OpE: ElementOp<E>,
2141>(
2142 dest: &mut StridedViewMut<D>,
2143 a: &StridedView<A, OpA>,
2144 b: &StridedView<B, OpB>,
2145 c: &StridedView<C, OpC>,
2146 e: &StridedView<E, OpE>,
2147 f: impl Fn(A, B, C, E) -> D + MaybeSync,
2148) -> Result<()> {
2149 ensure_same_shape(dest.dims(), a.dims())?;
2150 ensure_same_shape(dest.dims(), b.dims())?;
2151 ensure_same_shape(dest.dims(), c.dims())?;
2152 ensure_same_shape(dest.dims(), e.dims())?;
2153 let validated = validate_destination_layout(dest.dims(), dest.strides())?;
2154 zip_map4_into_validated(dest, a, b, c, e, f, validated)
2155}
2156
2157pub(crate) fn zip_map4_into_validated<
2158 D: Copy + MaybeSendSync,
2159 A: Copy + MaybeSendSync,
2160 B: Copy + MaybeSendSync,
2161 C: Copy + MaybeSendSync,
2162 E: Copy + MaybeSendSync,
2163 OpA: ElementOp<A>,
2164 OpB: ElementOp<B>,
2165 OpC: ElementOp<C>,
2166 OpE: ElementOp<E>,
2167>(
2168 dest: &mut StridedViewMut<D>,
2169 a: &StridedView<A, OpA>,
2170 b: &StridedView<B, OpB>,
2171 c: &StridedView<C, OpC>,
2172 e: &StridedView<E, OpE>,
2173 f: impl Fn(A, B, C, E) -> D + MaybeSync,
2174 _validated: ValidatedDestinationLayout,
2175) -> Result<()> {
2176 ensure_same_shape(dest.dims(), a.dims())?;
2177 ensure_same_shape(dest.dims(), b.dims())?;
2178 ensure_same_shape(dest.dims(), c.dims())?;
2179 ensure_same_shape(dest.dims(), e.dims())?;
2180 let dst_ptr = dest.as_mut_ptr();
2181 let a_ptr = a.ptr();
2182 let b_ptr = b.ptr();
2183 let c_ptr = c.ptr();
2184 let e_ptr = e.ptr();
2185
2186 let dst_dims = dest.dims();
2187 let dst_strides = dest.strides();
2188
2189 if sequential_contiguous_layout(
2190 dst_dims,
2191 &[
2192 dst_strides,
2193 a.strides(),
2194 b.strides(),
2195 c.strides(),
2196 e.strides(),
2197 ],
2198 )?
2199 .is_some()
2200 {
2201 let len = total_len(dst_dims)?;
2202 let dst = unsafe { std::slice::from_raw_parts_mut(dst_ptr, len) };
2203 let sa = unsafe { std::slice::from_raw_parts(a_ptr, len) };
2204 let sb = unsafe { std::slice::from_raw_parts(b_ptr, len) };
2205 let sc = unsafe { std::slice::from_raw_parts(c_ptr, len) };
2206 let se = unsafe { std::slice::from_raw_parts(e_ptr, len) };
2207 simd::dispatch_if_large(len, || {
2208 for i in 0..len {
2209 dst[i] = f(
2210 OpA::apply(sa[i]),
2211 OpB::apply(sb[i]),
2212 OpC::apply(sc[i]),
2213 OpE::apply(se[i]),
2214 );
2215 }
2216 });
2217 return Ok(());
2218 }
2219
2220 let strides_list: [&[isize]; 5] = [
2221 dst_strides,
2222 a.strides(),
2223 b.strides(),
2224 c.strides(),
2225 e.strides(),
2226 ];
2227 let elem_size = std::mem::size_of::<D>()
2228 .max(std::mem::size_of::<A>())
2229 .max(std::mem::size_of::<B>())
2230 .max(std::mem::size_of::<C>())
2231 .max(std::mem::size_of::<E>());
2232 let total = total_len(dst_dims)?;
2233
2234 let (fused_dims, ordered_strides, plan) = if total <= SMALL_TENSOR_THRESHOLD {
2236 build_plan_fused_small(dst_dims, &strides_list)
2237 } else {
2238 build_plan_fused(dst_dims, &strides_list, Some(0), elem_size)
2239 };
2240
2241 #[cfg(feature = "parallel")]
2242 {
2243 let total = total_len(&fused_dims)?;
2244 let nthreads = crate::execution_policy::rayon_threads();
2245 if total > MINTHREADLENGTH && nthreads > 1 {
2246 use crate::threading::SendPtr;
2247 let dst_send = SendPtr(dst_ptr);
2248 let a_send = SendPtr(a_ptr as *mut A);
2249 let b_send = SendPtr(b_ptr as *mut B);
2250 let c_send = SendPtr(c_ptr as *mut C);
2251 let e_send = SendPtr(e_ptr as *mut E);
2252
2253 let costs = compute_costs(&ordered_strides);
2254 let initial_offsets = vec![0isize; strides_list.len()];
2255 return mapreduce_threaded(
2256 &fused_dims,
2257 &plan.block,
2258 &ordered_strides,
2259 &initial_offsets,
2260 &costs,
2261 nthreads,
2262 0,
2263 1,
2264 &|dims, blocks, strides_list, offsets| {
2265 for_each_inner_block_with_offsets(
2266 dims,
2267 blocks,
2268 strides_list,
2269 offsets,
2270 |offsets, len, strides| {
2271 let dp = unsafe { dst_send.as_ptr().offset(offsets[0]) };
2272 let ap = unsafe { a_send.as_const().offset(offsets[1]) };
2273 let bp = unsafe { b_send.as_const().offset(offsets[2]) };
2274 let cp = unsafe { c_send.as_const().offset(offsets[3]) };
2275 let ep = unsafe { e_send.as_const().offset(offsets[4]) };
2276 unsafe {
2277 inner_loop_map4::<D, A, B, C, E, OpA, OpB, OpC, OpE>(
2278 dp, strides[0], ap, strides[1], bp, strides[2], cp, strides[3],
2279 ep, strides[4], len, &f,
2280 )
2281 };
2282 Ok(())
2283 },
2284 )
2285 },
2286 );
2287 }
2288 }
2289
2290 let initial_offsets = vec![0isize; ordered_strides.len()];
2291 for_each_inner_block_preordered(
2292 &fused_dims,
2293 &plan.block,
2294 &ordered_strides,
2295 &initial_offsets,
2296 |offsets, len, strides| {
2297 let dp = unsafe { dst_ptr.offset(offsets[0]) };
2298 let ap = unsafe { a_ptr.offset(offsets[1]) };
2299 let bp = unsafe { b_ptr.offset(offsets[2]) };
2300 let cp = unsafe { c_ptr.offset(offsets[3]) };
2301 let ep = unsafe { e_ptr.offset(offsets[4]) };
2302 unsafe {
2303 inner_loop_map4::<D, A, B, C, E, OpA, OpB, OpC, OpE>(
2304 dp, strides[0], ap, strides[1], bp, strides[2], cp, strides[3], ep, strides[4],
2305 len, &f,
2306 )
2307 };
2308 Ok(())
2309 },
2310 )
2311}
2312
2313#[cfg(test)]
2314#[path = "map_view/tests/scalar_branch_tests.rs"]
2315mod scalar_branch_tests;
2316
2317#[cfg(all(test, feature = "parallel"))]
2318#[path = "map_view/tests/tests.rs"]
2319mod tests;