Skip to main content

strided_basic/
map_view.rs

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