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 (((d, &a), &b), &c) in dst.iter_mut().zip(src_a).zip(src_b).zip(src_c) {
998 *d = f(OpA::apply(a), OpB::apply(b), OpC::apply(c));
999 }
1000 });
1001 } else {
1002 let mut dp = dp;
1003 let mut ap = ap;
1004 let mut bp = bp;
1005 let mut cp = cp;
1006 for _ in 0..len {
1007 *dp = f(OpA::apply(*ap), OpB::apply(*bp), OpC::apply(*cp));
1008 dp = dp.offset(ds);
1009 ap = ap.offset(a_s);
1010 bp = bp.offset(b_s);
1011 cp = cp.offset(c_s);
1012 }
1013 }
1014}
1015
1016#[inline(always)]
1018unsafe fn inner_loop_map4<
1019 D: Copy,
1020 A: Copy,
1021 B: Copy,
1022 C: Copy,
1023 E: Copy,
1024 OpA: ElementOp<A>,
1025 OpB: ElementOp<B>,
1026 OpC: ElementOp<C>,
1027 OpE: ElementOp<E>,
1028>(
1029 dp: *mut D,
1030 ds: isize,
1031 ap: *const A,
1032 a_s: isize,
1033 bp: *const B,
1034 b_s: isize,
1035 cp: *const C,
1036 c_s: isize,
1037 ep: *const E,
1038 e_s: isize,
1039 len: usize,
1040 f: &impl Fn(A, B, C, E) -> D,
1041) {
1042 if ds == 1 && a_s == 1 && b_s == 1 && c_s == 1 && e_s == 1 {
1043 let src_a = std::slice::from_raw_parts(ap, len);
1044 let src_b = std::slice::from_raw_parts(bp, len);
1045 let src_c = std::slice::from_raw_parts(cp, len);
1046 let src_e = std::slice::from_raw_parts(ep, len);
1047 let dst = std::slice::from_raw_parts_mut(dp, len);
1048 simd::dispatch_if_large(len, || {
1049 for i in 0..len {
1050 dst[i] = f(
1051 OpA::apply(src_a[i]),
1052 OpB::apply(src_b[i]),
1053 OpC::apply(src_c[i]),
1054 OpE::apply(src_e[i]),
1055 );
1056 }
1057 });
1058 } else {
1059 let mut dp = dp;
1060 let mut ap = ap;
1061 let mut bp = bp;
1062 let mut cp = cp;
1063 let mut ep = ep;
1064 for _ in 0..len {
1065 *dp = f(
1066 OpA::apply(*ap),
1067 OpB::apply(*bp),
1068 OpC::apply(*cp),
1069 OpE::apply(*ep),
1070 );
1071 dp = dp.offset(ds);
1072 ap = ap.offset(a_s);
1073 bp = bp.offset(b_s);
1074 cp = cp.offset(c_s);
1075 ep = ep.offset(e_s);
1076 }
1077 }
1078}
1079
1080pub fn map_into<D: Copy + MaybeSendSync, A: Copy + MaybeSendSync, Op: ElementOp<A>>(
1085 dest: &mut StridedViewMut<D>,
1086 src: &StridedView<A, Op>,
1087 f: impl Fn(A) -> D + MaybeSync,
1088) -> Result<()> {
1089 map_parts_into::<D, A, Op>(
1090 dest.as_mut_ptr(),
1091 dest.dims(),
1092 dest.strides(),
1093 src.ptr(),
1094 src.dims(),
1095 src.strides(),
1096 f,
1097 )
1098}
1099
1100pub(crate) fn map_into_validated<
1101 D: Copy + MaybeSendSync,
1102 A: Copy + MaybeSendSync,
1103 Op: ElementOp<A>,
1104>(
1105 dest: &mut StridedViewMut<D>,
1106 src: &StridedView<A, Op>,
1107 f: impl Fn(A) -> D + MaybeSync,
1108 validated: ValidatedDestinationLayout,
1109) -> Result<()> {
1110 ensure_same_shape(dest.dims(), src.dims())?;
1111 map_parts_into_validated::<D, A, Op>(
1112 dest.as_mut_ptr(),
1113 dest.dims(),
1114 dest.strides(),
1115 src.ptr(),
1116 src.strides(),
1117 f,
1118 validated,
1119 )
1120}
1121
1122pub(crate) fn map_raw_into<D: Copy + MaybeSendSync, A: Copy + MaybeSendSync, Op: ElementOp<A>>(
1123 dest: &mut crate::RawStridedMut<'_, D>,
1124 src: &crate::RawStridedRef<'_, A>,
1125 f: impl Fn(A) -> D + MaybeSync,
1126) -> Result<()> {
1127 map_parts_into::<D, A, Op>(
1128 dest.as_mut_ptr(),
1129 dest.dims(),
1130 dest.strides(),
1131 src.ptr(),
1132 src.dims(),
1133 src.strides(),
1134 f,
1135 )
1136}
1137
1138pub(crate) fn map_raw_into_validated<
1139 D: Copy + MaybeSendSync,
1140 A: Copy + MaybeSendSync,
1141 Op: ElementOp<A>,
1142>(
1143 dest: &mut crate::RawStridedMut<'_, D>,
1144 src: &crate::RawStridedRef<'_, A>,
1145 f: impl Fn(A) -> D + MaybeSync,
1146 validated: ValidatedDestinationLayout,
1147) -> Result<()> {
1148 ensure_same_shape(dest.dims(), src.dims())?;
1149 map_parts_into_validated::<D, A, Op>(
1150 dest.as_mut_ptr(),
1151 dest.dims(),
1152 dest.strides(),
1153 src.ptr(),
1154 src.strides(),
1155 f,
1156 validated,
1157 )
1158}
1159
1160#[allow(clippy::too_many_arguments)]
1161fn map_parts_into<D: Copy + MaybeSendSync, A: Copy + MaybeSendSync, Op: ElementOp<A>>(
1162 dst_ptr: *mut D,
1163 dst_dims: &[usize],
1164 dst_strides: &[isize],
1165 src_ptr: *const A,
1166 src_dims: &[usize],
1167 src_strides: &[isize],
1168 f: impl Fn(A) -> D + MaybeSync,
1169) -> Result<()> {
1170 ensure_same_shape(dst_dims, src_dims)?;
1171 let validated = validate_destination_layout(dst_dims, dst_strides)?;
1172 map_parts_into_validated::<D, A, Op>(
1173 dst_ptr,
1174 dst_dims,
1175 dst_strides,
1176 src_ptr,
1177 src_strides,
1178 f,
1179 validated,
1180 )
1181}
1182
1183#[allow(clippy::too_many_arguments)]
1184fn map_parts_into_validated<D: Copy + MaybeSendSync, A: Copy + MaybeSendSync, Op: ElementOp<A>>(
1185 dst_ptr: *mut D,
1186 dst_dims: &[usize],
1187 dst_strides: &[isize],
1188 src_ptr: *const A,
1189 src_strides: &[isize],
1190 f: impl Fn(A) -> D + MaybeSync,
1191 _validated: ValidatedDestinationLayout,
1192) -> Result<()> {
1193 if sequential_contiguous_layout(dst_dims, &[dst_strides, src_strides])?.is_some() {
1194 let len = total_len(dst_dims)?;
1195 let dst = unsafe { std::slice::from_raw_parts_mut(dst_ptr, len) };
1196 let src = unsafe { std::slice::from_raw_parts(src_ptr, len) };
1197 simd::dispatch_if_large(len, || {
1198 for i in 0..len {
1199 dst[i] = f(Op::apply(src[i]));
1200 }
1201 });
1202 return Ok(());
1203 }
1204
1205 let strides_list: [&[isize]; 2] = [dst_strides, src_strides];
1206 let elem_size = std::mem::size_of::<D>().max(std::mem::size_of::<A>());
1207 let total = total_len(dst_dims)?;
1208
1209 let (fused_dims, ordered_strides, plan) = if total <= SMALL_TENSOR_THRESHOLD {
1211 build_plan_fused_small(dst_dims, &strides_list)
1212 } else {
1213 build_plan_fused(dst_dims, &strides_list, Some(0), elem_size)
1214 };
1215
1216 #[cfg(feature = "parallel")]
1217 {
1218 let total = total_len(&fused_dims)?;
1219 let nthreads = crate::execution_policy::rayon_threads();
1220 if total > MINTHREADLENGTH && nthreads > 1 {
1221 use crate::threading::SendPtr;
1222 let dst_send = SendPtr(dst_ptr);
1223 let src_send = SendPtr(src_ptr as *mut A);
1224
1225 let costs = compute_costs(&ordered_strides);
1226 let initial_offsets = vec![0isize; strides_list.len()];
1227 return mapreduce_threaded(
1228 &fused_dims,
1229 &plan.block,
1230 &ordered_strides,
1231 &initial_offsets,
1232 &costs,
1233 nthreads,
1234 0,
1235 1,
1236 &|dims, blocks, strides_list, offsets| {
1237 for_each_inner_block_with_offsets(
1238 dims,
1239 blocks,
1240 strides_list,
1241 offsets,
1242 |offsets, len, strides| {
1243 let dp = unsafe { dst_send.as_ptr().offset(offsets[0]) };
1244 let sp = unsafe { src_send.as_const().offset(offsets[1]) };
1245 unsafe {
1246 inner_loop_map1::<D, A, Op>(dp, strides[0], sp, strides[1], len, &f)
1247 };
1248 Ok(())
1249 },
1250 )
1251 },
1252 );
1253 }
1254 }
1255
1256 let initial_offsets = vec![0isize; ordered_strides.len()];
1257 for_each_inner_block_preordered(
1258 &fused_dims,
1259 &plan.block,
1260 &ordered_strides,
1261 &initial_offsets,
1262 |offsets, len, strides| {
1263 let dp = unsafe { dst_ptr.offset(offsets[0]) };
1264 let sp = unsafe { src_ptr.offset(offsets[1]) };
1265 unsafe { inner_loop_map1::<D, A, Op>(dp, strides[0], sp, strides[1], len, &f) };
1266 Ok(())
1267 },
1268 )
1269}
1270
1271pub fn zip_map2_into<
1276 D: Copy + MaybeSendSync,
1277 A: Copy + MaybeSendSync,
1278 B: Copy + MaybeSendSync,
1279 OpA: ElementOp<A>,
1280 OpB: ElementOp<B>,
1281>(
1282 dest: &mut StridedViewMut<D>,
1283 a: &StridedView<A, OpA>,
1284 b: &StridedView<B, OpB>,
1285 f: impl Fn(A, B) -> D + MaybeSync,
1286) -> Result<()> {
1287 zip_map2_parts_into::<D, A, B, OpA, OpB>(
1288 dest.as_mut_ptr(),
1289 dest.dims(),
1290 dest.strides(),
1291 a.ptr(),
1292 a.dims(),
1293 a.strides(),
1294 b.ptr(),
1295 b.dims(),
1296 b.strides(),
1297 f,
1298 )
1299}
1300
1301pub(crate) fn zip_map2_into_validated<
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 validated: ValidatedDestinationLayout,
1313) -> Result<()> {
1314 zip_map2_parts_into_validated::<D, A, B, OpA, OpB>(
1315 dest.as_mut_ptr(),
1316 dest.dims(),
1317 dest.strides(),
1318 a.ptr(),
1319 a.strides(),
1320 b.ptr(),
1321 b.strides(),
1322 f,
1323 validated,
1324 )
1325}
1326
1327#[non_exhaustive]
1329#[derive(Clone, Copy, Debug, Eq, PartialEq)]
1330pub enum CompareOp {
1331 Eq,
1332 Lt,
1333 Le,
1334 Gt,
1335 Ge,
1336}
1337
1338pub fn compare_into<T, OpA, OpB>(
1349 dest: &mut StridedViewMut<bool>,
1350 a: &StridedView<T, OpA>,
1351 b: &StridedView<T, OpB>,
1352 op: CompareOp,
1353) -> Result<()>
1354where
1355 T: Copy + MaybeSendSync + PartialOrd,
1356 OpA: ElementOp<T>,
1357 OpB: ElementOp<T>,
1358{
1359 match op {
1360 CompareOp::Eq => zip_map2_into(dest, a, b, |lhs, rhs| lhs == rhs),
1361 CompareOp::Lt => zip_map2_into(dest, a, b, |lhs, rhs| lhs < rhs),
1362 CompareOp::Le => zip_map2_into(dest, a, b, |lhs, rhs| lhs <= rhs),
1363 CompareOp::Gt => zip_map2_into(dest, a, b, |lhs, rhs| lhs > rhs),
1364 CompareOp::Ge => zip_map2_into(dest, a, b, |lhs, rhs| lhs >= rhs),
1365 }
1366}
1367
1368pub fn compare_into_uninit<T, OpA, OpB>(
1387 dest: &mut StridedViewMut<MaybeUninit<bool>>,
1388 a: &StridedView<T, OpA>,
1389 b: &StridedView<T, OpB>,
1390 op: CompareOp,
1391) -> Result<()>
1392where
1393 T: Copy + MaybeSendSync + PartialOrd,
1394 OpA: ElementOp<T>,
1395 OpB: ElementOp<T>,
1396{
1397 ensure_same_shape(dest.dims(), a.dims())?;
1398 ensure_same_shape(dest.dims(), b.dims())?;
1399 let validated = validate_destination_layout(dest.dims(), dest.strides())?;
1400 validate_typed_no_overlap(dest, a, 0)?;
1401 validate_typed_no_overlap(dest, b, 1)?;
1402 match op {
1403 CompareOp::Eq => zip_map2_into_validated(
1404 dest,
1405 a,
1406 b,
1407 |lhs, rhs| MaybeUninit::new(lhs == rhs),
1408 validated,
1409 ),
1410 CompareOp::Lt => zip_map2_into_validated(
1411 dest,
1412 a,
1413 b,
1414 |lhs, rhs| MaybeUninit::new(lhs < rhs),
1415 validated,
1416 ),
1417 CompareOp::Le => zip_map2_into_validated(
1418 dest,
1419 a,
1420 b,
1421 |lhs, rhs| MaybeUninit::new(lhs <= rhs),
1422 validated,
1423 ),
1424 CompareOp::Gt => zip_map2_into_validated(
1425 dest,
1426 a,
1427 b,
1428 |lhs, rhs| MaybeUninit::new(lhs > rhs),
1429 validated,
1430 ),
1431 CompareOp::Ge => zip_map2_into_validated(
1432 dest,
1433 a,
1434 b,
1435 |lhs, rhs| MaybeUninit::new(lhs >= rhs),
1436 validated,
1437 ),
1438 }
1439}
1440
1441pub(crate) fn zip_map2_raw_into_validated<
1442 D: Copy + MaybeSendSync,
1443 A: Copy + MaybeSendSync,
1444 B: Copy + MaybeSendSync,
1445 OpA: ElementOp<A>,
1446 OpB: ElementOp<B>,
1447>(
1448 dest: &mut crate::RawStridedMut<'_, D>,
1449 a: &crate::RawStridedRef<'_, A>,
1450 b: &crate::RawStridedRef<'_, B>,
1451 f: impl Fn(A, B) -> D + MaybeSync,
1452 validated: ValidatedDestinationLayout,
1453) -> Result<()> {
1454 ensure_same_shape(dest.dims(), a.dims())?;
1455 ensure_same_shape(dest.dims(), b.dims())?;
1456 zip_map2_parts_into_validated::<D, A, B, OpA, OpB>(
1457 dest.as_mut_ptr(),
1458 dest.dims(),
1459 dest.strides(),
1460 a.ptr(),
1461 a.strides(),
1462 b.ptr(),
1463 b.strides(),
1464 f,
1465 validated,
1466 )
1467}
1468
1469#[allow(clippy::too_many_arguments)]
1470fn zip_map2_parts_into<
1471 D: Copy + MaybeSendSync,
1472 A: Copy + MaybeSendSync,
1473 B: Copy + MaybeSendSync,
1474 OpA: ElementOp<A>,
1475 OpB: ElementOp<B>,
1476>(
1477 dst_ptr: *mut D,
1478 dst_dims: &[usize],
1479 dst_strides: &[isize],
1480 a_ptr: *const A,
1481 a_dims: &[usize],
1482 a_strides: &[isize],
1483 b_ptr: *const B,
1484 b_dims: &[usize],
1485 b_strides: &[isize],
1486 f: impl Fn(A, B) -> D + MaybeSync,
1487) -> Result<()> {
1488 ensure_same_shape(dst_dims, a_dims)?;
1489 ensure_same_shape(dst_dims, b_dims)?;
1490 let validated = validate_destination_layout(dst_dims, dst_strides)?;
1491 zip_map2_parts_into_validated::<D, A, B, OpA, OpB>(
1492 dst_ptr,
1493 dst_dims,
1494 dst_strides,
1495 a_ptr,
1496 a_strides,
1497 b_ptr,
1498 b_strides,
1499 f,
1500 validated,
1501 )
1502}
1503
1504#[allow(clippy::too_many_arguments)]
1505fn zip_map2_parts_into_validated<
1506 D: Copy + MaybeSendSync,
1507 A: Copy + MaybeSendSync,
1508 B: Copy + MaybeSendSync,
1509 OpA: ElementOp<A>,
1510 OpB: ElementOp<B>,
1511>(
1512 dst_ptr: *mut D,
1513 dst_dims: &[usize],
1514 dst_strides: &[isize],
1515 a_ptr: *const A,
1516 a_strides: &[isize],
1517 b_ptr: *const B,
1518 b_strides: &[isize],
1519 f: impl Fn(A, B) -> D + MaybeSync,
1520 _validated: ValidatedDestinationLayout,
1521) -> Result<()> {
1522 if sequential_contiguous_layout(dst_dims, &[dst_strides, a_strides, b_strides])?.is_some() {
1523 let len = total_len(dst_dims)?;
1524 let dst = unsafe { std::slice::from_raw_parts_mut(dst_ptr, len) };
1525 let sa = unsafe { std::slice::from_raw_parts(a_ptr, len) };
1526 let sb = unsafe { std::slice::from_raw_parts(b_ptr, len) };
1527 simd::dispatch_if_large(len, || {
1528 for i in 0..len {
1529 dst[i] = f(OpA::apply(sa[i]), OpB::apply(sb[i]));
1530 }
1531 });
1532 return Ok(());
1533 }
1534
1535 let strides_list: [&[isize]; 3] = [dst_strides, a_strides, b_strides];
1536 let elem_size = std::mem::size_of::<D>()
1537 .max(std::mem::size_of::<A>())
1538 .max(std::mem::size_of::<B>());
1539 let total = total_len(dst_dims)?;
1540
1541 let (fused_dims, ordered_strides, plan) = if total <= SMALL_TENSOR_THRESHOLD {
1543 build_plan_fused_small(dst_dims, &strides_list)
1544 } else {
1545 build_plan_fused(dst_dims, &strides_list, Some(0), elem_size)
1546 };
1547
1548 #[cfg(feature = "parallel")]
1549 {
1550 let total = total_len(&fused_dims)?;
1551 let nthreads = crate::execution_policy::rayon_threads();
1552 if total > MINTHREADLENGTH && nthreads > 1 {
1553 use crate::threading::SendPtr;
1554 let dst_send = SendPtr(dst_ptr);
1555 let a_send = SendPtr(a_ptr as *mut A);
1556 let b_send = SendPtr(b_ptr as *mut B);
1557
1558 let costs = compute_costs(&ordered_strides);
1559 let initial_offsets = vec![0isize; strides_list.len()];
1560 return mapreduce_threaded(
1561 &fused_dims,
1562 &plan.block,
1563 &ordered_strides,
1564 &initial_offsets,
1565 &costs,
1566 nthreads,
1567 0,
1568 1,
1569 &|dims, blocks, strides_list, offsets| {
1570 for_each_inner_block_with_offsets(
1571 dims,
1572 blocks,
1573 strides_list,
1574 offsets,
1575 |offsets, len, strides| {
1576 let dp = unsafe { dst_send.as_ptr().offset(offsets[0]) };
1577 let ap = unsafe { a_send.as_const().offset(offsets[1]) };
1578 let bp = unsafe { b_send.as_const().offset(offsets[2]) };
1579 unsafe {
1580 inner_loop_map2::<D, A, B, OpA, OpB>(
1581 dp, strides[0], ap, strides[1], bp, strides[2], len, &f,
1582 )
1583 };
1584 Ok(())
1585 },
1586 )
1587 },
1588 );
1589 }
1590 }
1591
1592 let initial_offsets = vec![0isize; ordered_strides.len()];
1593 for_each_inner_block_preordered(
1594 &fused_dims,
1595 &plan.block,
1596 &ordered_strides,
1597 &initial_offsets,
1598 |offsets, len, strides| {
1599 let dp = unsafe { dst_ptr.offset(offsets[0]) };
1600 let ap = unsafe { a_ptr.offset(offsets[1]) };
1601 let bp = unsafe { b_ptr.offset(offsets[2]) };
1602 unsafe {
1603 inner_loop_map2::<D, A, B, OpA, OpB>(
1604 dp, strides[0], ap, strides[1], bp, strides[2], len, &f,
1605 )
1606 };
1607 Ok(())
1608 },
1609 )
1610}
1611
1612fn mul_identity_into_raw<
1613 O: MulOutput<D>,
1614 D: Copy + MaybeSendSync + 'static,
1615 A: Copy + MaybeSendSync + Mul<B, Output = D> + 'static,
1616 B: Copy + MaybeSendSync + 'static,
1617>(
1618 dst_ptr: *mut O::Slot,
1619 dst_dims: &[usize],
1620 dst_strides: &[isize],
1621 a_ptr: *const A,
1622 a_strides: &[isize],
1623 b_ptr: *const B,
1624 b_strides: &[isize],
1625 _validated: ValidatedDestinationLayout,
1626) -> Result<()> {
1627 debug_assert_eq!(dst_dims.len(), a_strides.len());
1628 debug_assert_eq!(dst_dims.len(), b_strides.len());
1629
1630 if sequential_contiguous_layout(dst_dims, &[dst_strides, a_strides, b_strides])?.is_some() {
1631 let len = total_len(dst_dims)?;
1632 let sa = unsafe { std::slice::from_raw_parts(a_ptr, len) };
1633 let sb = unsafe { std::slice::from_raw_parts(b_ptr, len) };
1634 if unsafe { O::try_contiguous(dst_ptr, len, sa, sb) } {
1635 return Ok(());
1636 }
1637 for i in 0..len {
1638 unsafe { O::write(dst_ptr.add(i), multiply_value(sa[i], sb[i])) };
1639 }
1640 return Ok(());
1641 }
1642
1643 let strides_list: [&[isize]; 3] = [dst_strides, a_strides, b_strides];
1644 let elem_size = std::mem::size_of::<D>()
1645 .max(std::mem::size_of::<A>())
1646 .max(std::mem::size_of::<B>());
1647 let total = total_len(dst_dims)?;
1648
1649 if try_contiguous_range_mul::<O, D, A, B>(
1650 dst_ptr,
1651 dst_dims,
1652 dst_strides,
1653 a_ptr,
1654 a_strides,
1655 b_ptr,
1656 b_strides,
1657 ) {
1658 return Ok(());
1659 }
1660
1661 let (fused_dims, ordered_strides, plan) = if total <= SMALL_TENSOR_THRESHOLD {
1662 build_plan_fused_small(dst_dims, &strides_list)
1663 } else {
1664 build_plan_fused(dst_dims, &strides_list, Some(0), elem_size)
1665 };
1666
1667 #[cfg(feature = "parallel")]
1668 {
1669 let total = total_len(&fused_dims)?;
1670 let nthreads = crate::execution_policy::rayon_threads();
1671 if total > MINTHREADLENGTH && nthreads > 1 {
1672 use crate::threading::SendPtr;
1673 let dst_send = SendPtr(dst_ptr);
1674 let a_send = SendPtr(a_ptr as *mut A);
1675 let b_send = SendPtr(b_ptr as *mut B);
1676
1677 let costs = compute_costs(&ordered_strides);
1678 let initial_offsets = vec![0isize; strides_list.len()];
1679 return mapreduce_threaded(
1680 &fused_dims,
1681 &plan.block,
1682 &ordered_strides,
1683 &initial_offsets,
1684 &costs,
1685 nthreads,
1686 0,
1687 1,
1688 &|dims, blocks, strides_list, offsets| {
1689 for_each_inner_block_with_offsets(
1690 dims,
1691 blocks,
1692 strides_list,
1693 offsets,
1694 |offsets, len, strides| {
1695 let dp = unsafe { dst_send.as_ptr().offset(offsets[0]) };
1696 let ap = unsafe { a_send.as_const().offset(offsets[1]) };
1697 let bp = unsafe { b_send.as_const().offset(offsets[2]) };
1698 unsafe {
1699 inner_loop_mul2::<O, D, A, B>(
1700 dp, strides[0], ap, strides[1], bp, strides[2], len,
1701 )
1702 };
1703 Ok(())
1704 },
1705 )
1706 },
1707 );
1708 }
1709 }
1710
1711 let initial_offsets = vec![0isize; ordered_strides.len()];
1712 for_each_inner_block_preordered(
1713 &fused_dims,
1714 &plan.block,
1715 &ordered_strides,
1716 &initial_offsets,
1717 |offsets, len, strides| {
1718 let dp = unsafe { dst_ptr.offset(offsets[0]) };
1719 let ap = unsafe { a_ptr.offset(offsets[1]) };
1720 let bp = unsafe { b_ptr.offset(offsets[2]) };
1721 unsafe {
1722 inner_loop_mul2::<O, D, A, B>(dp, strides[0], ap, strides[1], bp, strides[2], len)
1723 };
1724 Ok(())
1725 },
1726 )
1727}
1728
1729pub fn mul_into<
1734 D: Copy + MaybeSendSync + 'static,
1735 A: Copy + Mul<B, Output = D> + MaybeSendSync + 'static,
1736 B: Copy + MaybeSendSync + 'static,
1737 OpA: ElementOp<A>,
1738 OpB: ElementOp<B>,
1739>(
1740 dest: &mut StridedViewMut<D>,
1741 a: &StridedView<A, OpA>,
1742 b: &StridedView<B, OpB>,
1743) -> Result<()> {
1744 ensure_same_shape(dest.dims(), a.dims())?;
1745 ensure_same_shape(dest.dims(), b.dims())?;
1746 let validated = validate_destination_layout(dest.dims(), dest.strides())?;
1747
1748 if OpA::IS_IDENTITY && OpB::IS_IDENTITY {
1749 return mul_identity_into_raw::<InitializedOutput, D, A, B>(
1750 dest.as_mut_ptr(),
1751 dest.dims(),
1752 dest.strides(),
1753 a.ptr(),
1754 a.strides(),
1755 b.ptr(),
1756 b.strides(),
1757 validated,
1758 );
1759 }
1760
1761 zip_map2_into_validated(dest, a, b, multiply_value, validated)
1762}
1763
1764pub fn mul_into_uninit<
1783 D: Copy + MaybeSendSync + 'static,
1784 A: Copy + Mul<B, Output = D> + MaybeSendSync + 'static,
1785 B: Copy + MaybeSendSync + 'static,
1786 OpA: ElementOp<A>,
1787 OpB: ElementOp<B>,
1788>(
1789 dest: &mut StridedViewMut<MaybeUninit<D>>,
1790 a: &StridedView<A, OpA>,
1791 b: &StridedView<B, OpB>,
1792) -> Result<()> {
1793 ensure_same_shape(dest.dims(), a.dims())?;
1794 ensure_same_shape(dest.dims(), b.dims())?;
1795 let validated = validate_destination_layout(dest.dims(), dest.strides())?;
1796 validate_typed_no_overlap(dest, a, 0)?;
1797 validate_typed_no_overlap(dest, b, 1)?;
1798
1799 if OpA::IS_IDENTITY && OpB::IS_IDENTITY {
1800 return mul_identity_into_raw::<UninitializedOutput, D, A, B>(
1801 dest.as_mut_ptr(),
1802 dest.dims(),
1803 dest.strides(),
1804 a.ptr(),
1805 a.strides(),
1806 b.ptr(),
1807 b.strides(),
1808 validated,
1809 );
1810 }
1811
1812 zip_map2_into_validated(
1813 dest,
1814 a,
1815 b,
1816 |lhs, rhs| MaybeUninit::new(multiply_value(lhs, rhs)),
1817 validated,
1818 )
1819}
1820
1821fn broadcast_strides_for_axes(
1822 source_dims: &[usize],
1823 source_strides: &[isize],
1824 target_dims: &[usize],
1825 axes: &[usize],
1826) -> Result<AxisVec<isize>> {
1827 if source_dims.len() != axes.len() {
1828 return Err(StridedError::RankMismatch(source_dims.len(), axes.len()));
1829 }
1830 debug_assert_eq!(source_dims.len(), source_strides.len());
1831
1832 let mut seen = AxisVec::<bool>::new();
1833 seen.resize(target_dims.len(), false);
1834 let mut strides = AxisVec::<isize>::new();
1835 strides.resize(target_dims.len(), 0);
1836 for (src_axis, &dst_axis) in axes.iter().enumerate() {
1837 if dst_axis >= target_dims.len() {
1838 return Err(StridedError::InvalidAxis {
1839 axis: dst_axis,
1840 rank: target_dims.len(),
1841 });
1842 }
1843 if seen[dst_axis] {
1844 return Err(StridedError::InvalidAxis {
1845 axis: dst_axis,
1846 rank: target_dims.len(),
1847 });
1848 }
1849 seen[dst_axis] = true;
1850
1851 let source_dim = source_dims[src_axis];
1852 let target_dim = target_dims[dst_axis];
1853 if source_dim != target_dim && source_dim != 1 {
1854 return Err(StridedError::ShapeMismatch(
1855 source_dims.to_vec(),
1856 target_dims.to_vec(),
1857 ));
1858 }
1859 if source_dim == target_dim {
1860 strides[dst_axis] = source_strides[src_axis];
1861 }
1862 }
1863
1864 Ok(strides)
1865}
1866
1867fn broadcast_view_with_strides<'a, T, Op: ElementOp<T>>(
1868 view: &StridedView<'a, T, Op>,
1869 target_dims: &[usize],
1870 strides: &[isize],
1871) -> StridedView<'a, T, Op> {
1872 unsafe { StridedView::new_unchecked(view.data(), target_dims, strides, view.offset()) }
1873}
1874
1875pub fn broadcast_mul_into<
1880 D: Copy + MaybeSendSync + 'static,
1881 A: Copy + Mul<B, Output = D> + MaybeSendSync + 'static,
1882 B: Copy + MaybeSendSync + 'static,
1883 OpA: ElementOp<A>,
1884 OpB: ElementOp<B>,
1885>(
1886 dest: &mut StridedViewMut<D>,
1887 a: &StridedView<A, OpA>,
1888 a_axes: &[usize],
1889 b: &StridedView<B, OpB>,
1890 b_axes: &[usize],
1891) -> Result<()> {
1892 let validated = validate_destination_layout(dest.dims(), dest.strides())?;
1893 let a_strides = broadcast_strides_for_axes(a.dims(), a.strides(), dest.dims(), a_axes)?;
1894 let b_strides = broadcast_strides_for_axes(b.dims(), b.strides(), dest.dims(), b_axes)?;
1895
1896 if OpA::IS_IDENTITY && OpB::IS_IDENTITY {
1897 return mul_identity_into_raw::<InitializedOutput, D, A, B>(
1898 dest.as_mut_ptr(),
1899 dest.dims(),
1900 dest.strides(),
1901 a.ptr(),
1902 &a_strides,
1903 b.ptr(),
1904 &b_strides,
1905 validated,
1906 );
1907 }
1908
1909 let a = broadcast_view_with_strides(a, dest.dims(), &a_strides);
1910 let b = broadcast_view_with_strides(b, dest.dims(), &b_strides);
1911 zip_map2_into_validated(dest, &a, &b, multiply_value, validated)
1912}
1913
1914pub fn broadcast_mul_into_uninit<
1933 D: Copy + MaybeSendSync + 'static,
1934 A: Copy + Mul<B, Output = D> + MaybeSendSync + 'static,
1935 B: Copy + MaybeSendSync + 'static,
1936 OpA: ElementOp<A>,
1937 OpB: ElementOp<B>,
1938>(
1939 dest: &mut StridedViewMut<MaybeUninit<D>>,
1940 a: &StridedView<A, OpA>,
1941 a_axes: &[usize],
1942 b: &StridedView<B, OpB>,
1943 b_axes: &[usize],
1944) -> Result<()> {
1945 let validated = validate_destination_layout(dest.dims(), dest.strides())?;
1946 let a_strides = broadcast_strides_for_axes(a.dims(), a.strides(), dest.dims(), a_axes)?;
1947 let b_strides = broadcast_strides_for_axes(b.dims(), b.strides(), dest.dims(), b_axes)?;
1948 validate_typed_no_overlap(dest, a, 0)?;
1949 validate_typed_no_overlap(dest, b, 1)?;
1950
1951 if OpA::IS_IDENTITY && OpB::IS_IDENTITY {
1952 return mul_identity_into_raw::<UninitializedOutput, D, A, B>(
1953 dest.as_mut_ptr(),
1954 dest.dims(),
1955 dest.strides(),
1956 a.ptr(),
1957 &a_strides,
1958 b.ptr(),
1959 &b_strides,
1960 validated,
1961 );
1962 }
1963
1964 let a = broadcast_view_with_strides(a, dest.dims(), &a_strides);
1965 let b = broadcast_view_with_strides(b, dest.dims(), &b_strides);
1966 zip_map2_into_validated(
1967 dest,
1968 &a,
1969 &b,
1970 |lhs, rhs| MaybeUninit::new(multiply_value(lhs, rhs)),
1971 validated,
1972 )
1973}
1974
1975pub fn zip_map3_into<
1977 D: Copy + MaybeSendSync,
1978 A: Copy + MaybeSendSync,
1979 B: Copy + MaybeSendSync,
1980 C: Copy + MaybeSendSync,
1981 OpA: ElementOp<A>,
1982 OpB: ElementOp<B>,
1983 OpC: ElementOp<C>,
1984>(
1985 dest: &mut StridedViewMut<D>,
1986 a: &StridedView<A, OpA>,
1987 b: &StridedView<B, OpB>,
1988 c: &StridedView<C, OpC>,
1989 f: impl Fn(A, B, C) -> D + MaybeSync,
1990) -> Result<()> {
1991 ensure_same_shape(dest.dims(), a.dims())?;
1992 ensure_same_shape(dest.dims(), b.dims())?;
1993 ensure_same_shape(dest.dims(), c.dims())?;
1994 let validated = validate_destination_layout(dest.dims(), dest.strides())?;
1995 zip_map3_into_validated(dest, a, b, c, f, validated)
1996}
1997
1998pub(crate) fn zip_map3_into_validated<
1999 D: Copy + MaybeSendSync,
2000 A: Copy + MaybeSendSync,
2001 B: Copy + MaybeSendSync,
2002 C: Copy + MaybeSendSync,
2003 OpA: ElementOp<A>,
2004 OpB: ElementOp<B>,
2005 OpC: ElementOp<C>,
2006>(
2007 dest: &mut StridedViewMut<D>,
2008 a: &StridedView<A, OpA>,
2009 b: &StridedView<B, OpB>,
2010 c: &StridedView<C, OpC>,
2011 f: impl Fn(A, B, C) -> D + MaybeSync,
2012 _validated: ValidatedDestinationLayout,
2013) -> Result<()> {
2014 ensure_same_shape(dest.dims(), a.dims())?;
2015 ensure_same_shape(dest.dims(), b.dims())?;
2016 ensure_same_shape(dest.dims(), c.dims())?;
2017 let dst_ptr = dest.as_mut_ptr();
2018 let a_ptr = a.ptr();
2019 let b_ptr = b.ptr();
2020 let c_ptr = c.ptr();
2021
2022 let dst_dims = dest.dims();
2023 let dst_strides = dest.strides();
2024
2025 if sequential_contiguous_layout(
2026 dst_dims,
2027 &[dst_strides, a.strides(), b.strides(), c.strides()],
2028 )?
2029 .is_some()
2030 {
2031 let len = total_len(dst_dims)?;
2032 let dst = unsafe { std::slice::from_raw_parts_mut(dst_ptr, len) };
2033 let sa = unsafe { std::slice::from_raw_parts(a_ptr, len) };
2034 let sb = unsafe { std::slice::from_raw_parts(b_ptr, len) };
2035 let sc = unsafe { std::slice::from_raw_parts(c_ptr, len) };
2036 simd::dispatch_if_large(len, || {
2037 for (((d, &a), &b), &c) in dst.iter_mut().zip(sa).zip(sb).zip(sc) {
2038 *d = f(OpA::apply(a), OpB::apply(b), OpC::apply(c));
2039 }
2040 });
2041 return Ok(());
2042 }
2043
2044 let strides_list: [&[isize]; 4] = [dst_strides, a.strides(), b.strides(), c.strides()];
2045 let elem_size = std::mem::size_of::<D>()
2046 .max(std::mem::size_of::<A>())
2047 .max(std::mem::size_of::<B>())
2048 .max(std::mem::size_of::<C>());
2049 let total = total_len(dst_dims)?;
2050
2051 let (fused_dims, ordered_strides, plan) = if total <= SMALL_TENSOR_THRESHOLD {
2053 build_plan_fused_small(dst_dims, &strides_list)
2054 } else {
2055 build_plan_fused(dst_dims, &strides_list, Some(0), elem_size)
2056 };
2057
2058 #[cfg(feature = "parallel")]
2059 {
2060 let total = total_len(&fused_dims)?;
2061 let nthreads = crate::execution_policy::rayon_threads();
2062 if total > MINTHREADLENGTH && nthreads > 1 {
2063 use crate::threading::SendPtr;
2064 let dst_send = SendPtr(dst_ptr);
2065 let a_send = SendPtr(a_ptr as *mut A);
2066 let b_send = SendPtr(b_ptr as *mut B);
2067 let c_send = SendPtr(c_ptr as *mut C);
2068
2069 let costs = compute_costs(&ordered_strides);
2070 let initial_offsets = vec![0isize; strides_list.len()];
2071 return mapreduce_threaded(
2072 &fused_dims,
2073 &plan.block,
2074 &ordered_strides,
2075 &initial_offsets,
2076 &costs,
2077 nthreads,
2078 0,
2079 1,
2080 &|dims, blocks, strides_list, offsets| {
2081 for_each_inner_block_with_offsets(
2082 dims,
2083 blocks,
2084 strides_list,
2085 offsets,
2086 |offsets, len, strides| {
2087 let dp = unsafe { dst_send.as_ptr().offset(offsets[0]) };
2088 let ap = unsafe { a_send.as_const().offset(offsets[1]) };
2089 let bp = unsafe { b_send.as_const().offset(offsets[2]) };
2090 let cp = unsafe { c_send.as_const().offset(offsets[3]) };
2091 unsafe {
2092 inner_loop_map3::<D, A, B, C, OpA, OpB, OpC>(
2093 dp, strides[0], ap, strides[1], bp, strides[2], cp, strides[3],
2094 len, &f,
2095 )
2096 };
2097 Ok(())
2098 },
2099 )
2100 },
2101 );
2102 }
2103 }
2104
2105 let initial_offsets = vec![0isize; ordered_strides.len()];
2106 for_each_inner_block_preordered(
2107 &fused_dims,
2108 &plan.block,
2109 &ordered_strides,
2110 &initial_offsets,
2111 |offsets, len, strides| {
2112 let dp = unsafe { dst_ptr.offset(offsets[0]) };
2113 let ap = unsafe { a_ptr.offset(offsets[1]) };
2114 let bp = unsafe { b_ptr.offset(offsets[2]) };
2115 let cp = unsafe { c_ptr.offset(offsets[3]) };
2116 unsafe {
2117 inner_loop_map3::<D, A, B, C, OpA, OpB, OpC>(
2118 dp, strides[0], ap, strides[1], bp, strides[2], cp, strides[3], len, &f,
2119 )
2120 };
2121 Ok(())
2122 },
2123 )
2124}
2125
2126pub fn zip_map4_into<
2128 D: Copy + MaybeSendSync,
2129 A: Copy + MaybeSendSync,
2130 B: Copy + MaybeSendSync,
2131 C: Copy + MaybeSendSync,
2132 E: Copy + MaybeSendSync,
2133 OpA: ElementOp<A>,
2134 OpB: ElementOp<B>,
2135 OpC: ElementOp<C>,
2136 OpE: ElementOp<E>,
2137>(
2138 dest: &mut StridedViewMut<D>,
2139 a: &StridedView<A, OpA>,
2140 b: &StridedView<B, OpB>,
2141 c: &StridedView<C, OpC>,
2142 e: &StridedView<E, OpE>,
2143 f: impl Fn(A, B, C, E) -> D + MaybeSync,
2144) -> Result<()> {
2145 ensure_same_shape(dest.dims(), a.dims())?;
2146 ensure_same_shape(dest.dims(), b.dims())?;
2147 ensure_same_shape(dest.dims(), c.dims())?;
2148 ensure_same_shape(dest.dims(), e.dims())?;
2149 let validated = validate_destination_layout(dest.dims(), dest.strides())?;
2150 zip_map4_into_validated(dest, a, b, c, e, f, validated)
2151}
2152
2153pub(crate) fn zip_map4_into_validated<
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 _validated: ValidatedDestinationLayout,
2171) -> Result<()> {
2172 ensure_same_shape(dest.dims(), a.dims())?;
2173 ensure_same_shape(dest.dims(), b.dims())?;
2174 ensure_same_shape(dest.dims(), c.dims())?;
2175 ensure_same_shape(dest.dims(), e.dims())?;
2176 let dst_ptr = dest.as_mut_ptr();
2177 let a_ptr = a.ptr();
2178 let b_ptr = b.ptr();
2179 let c_ptr = c.ptr();
2180 let e_ptr = e.ptr();
2181
2182 let dst_dims = dest.dims();
2183 let dst_strides = dest.strides();
2184
2185 if sequential_contiguous_layout(
2186 dst_dims,
2187 &[
2188 dst_strides,
2189 a.strides(),
2190 b.strides(),
2191 c.strides(),
2192 e.strides(),
2193 ],
2194 )?
2195 .is_some()
2196 {
2197 let len = total_len(dst_dims)?;
2198 let dst = unsafe { std::slice::from_raw_parts_mut(dst_ptr, len) };
2199 let sa = unsafe { std::slice::from_raw_parts(a_ptr, len) };
2200 let sb = unsafe { std::slice::from_raw_parts(b_ptr, len) };
2201 let sc = unsafe { std::slice::from_raw_parts(c_ptr, len) };
2202 let se = unsafe { std::slice::from_raw_parts(e_ptr, len) };
2203 simd::dispatch_if_large(len, || {
2204 for i in 0..len {
2205 dst[i] = f(
2206 OpA::apply(sa[i]),
2207 OpB::apply(sb[i]),
2208 OpC::apply(sc[i]),
2209 OpE::apply(se[i]),
2210 );
2211 }
2212 });
2213 return Ok(());
2214 }
2215
2216 let strides_list: [&[isize]; 5] = [
2217 dst_strides,
2218 a.strides(),
2219 b.strides(),
2220 c.strides(),
2221 e.strides(),
2222 ];
2223 let elem_size = std::mem::size_of::<D>()
2224 .max(std::mem::size_of::<A>())
2225 .max(std::mem::size_of::<B>())
2226 .max(std::mem::size_of::<C>())
2227 .max(std::mem::size_of::<E>());
2228 let total = total_len(dst_dims)?;
2229
2230 let (fused_dims, ordered_strides, plan) = if total <= SMALL_TENSOR_THRESHOLD {
2232 build_plan_fused_small(dst_dims, &strides_list)
2233 } else {
2234 build_plan_fused(dst_dims, &strides_list, Some(0), elem_size)
2235 };
2236
2237 #[cfg(feature = "parallel")]
2238 {
2239 let total = total_len(&fused_dims)?;
2240 let nthreads = crate::execution_policy::rayon_threads();
2241 if total > MINTHREADLENGTH && nthreads > 1 {
2242 use crate::threading::SendPtr;
2243 let dst_send = SendPtr(dst_ptr);
2244 let a_send = SendPtr(a_ptr as *mut A);
2245 let b_send = SendPtr(b_ptr as *mut B);
2246 let c_send = SendPtr(c_ptr as *mut C);
2247 let e_send = SendPtr(e_ptr as *mut E);
2248
2249 let costs = compute_costs(&ordered_strides);
2250 let initial_offsets = vec![0isize; strides_list.len()];
2251 return mapreduce_threaded(
2252 &fused_dims,
2253 &plan.block,
2254 &ordered_strides,
2255 &initial_offsets,
2256 &costs,
2257 nthreads,
2258 0,
2259 1,
2260 &|dims, blocks, strides_list, offsets| {
2261 for_each_inner_block_with_offsets(
2262 dims,
2263 blocks,
2264 strides_list,
2265 offsets,
2266 |offsets, len, strides| {
2267 let dp = unsafe { dst_send.as_ptr().offset(offsets[0]) };
2268 let ap = unsafe { a_send.as_const().offset(offsets[1]) };
2269 let bp = unsafe { b_send.as_const().offset(offsets[2]) };
2270 let cp = unsafe { c_send.as_const().offset(offsets[3]) };
2271 let ep = unsafe { e_send.as_const().offset(offsets[4]) };
2272 unsafe {
2273 inner_loop_map4::<D, A, B, C, E, OpA, OpB, OpC, OpE>(
2274 dp, strides[0], ap, strides[1], bp, strides[2], cp, strides[3],
2275 ep, strides[4], len, &f,
2276 )
2277 };
2278 Ok(())
2279 },
2280 )
2281 },
2282 );
2283 }
2284 }
2285
2286 let initial_offsets = vec![0isize; ordered_strides.len()];
2287 for_each_inner_block_preordered(
2288 &fused_dims,
2289 &plan.block,
2290 &ordered_strides,
2291 &initial_offsets,
2292 |offsets, len, strides| {
2293 let dp = unsafe { dst_ptr.offset(offsets[0]) };
2294 let ap = unsafe { a_ptr.offset(offsets[1]) };
2295 let bp = unsafe { b_ptr.offset(offsets[2]) };
2296 let cp = unsafe { c_ptr.offset(offsets[3]) };
2297 let ep = unsafe { e_ptr.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], ep, strides[4],
2301 len, &f,
2302 )
2303 };
2304 Ok(())
2305 },
2306 )
2307}
2308
2309#[cfg(test)]
2310#[path = "map_view/tests/scalar_branch_tests.rs"]
2311mod scalar_branch_tests;
2312
2313#[cfg(all(test, feature = "parallel"))]
2314#[path = "map_view/tests/tests.rs"]
2315mod tests;