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