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// INVARIANT: validated layouts keep every dereferenced cursor in bounds.
170// Pointwise leaves use wrapping advances: the unused final cursor may lie
171// outside the allocation for negative or gapped strides.
172/// Unary inner loop: `dest[i] = f(Op::apply(src[i]))` for `len` elements.
173#[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/// Binary inner loop: `dest[i] = f(OpA::apply(a[i]), OpB::apply(b[i]))`.
202#[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
282/// Binary multiplication inner loop for identity element ops.
283trait 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        // SIMD stores write complete values directly through raw pointers; no
337        // initialized reference is formed over the uninitialized destination.
338        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        // The exact TypeId checks prove identical layout and validity.
356        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        // The exact TypeId checks prove identical layout and validity.
366        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    /// Whether running the plan row by row, without cache blocking, is sound
459    /// for performance.
460    ///
461    /// The destination is walked contiguously, so a unit-stride or broadcast
462    /// (stride 0) input streams along with it. An input read at any other
463    /// stride along the fast axis (a transposed operand) is swept across the
464    /// whole fast axis once per row, touching a new cache line per element and
465    /// again on the next row; the blocked kernel tiles that access. The one
466    /// exception is the SIMD transposed-scalar tile, which blocks itself.
467    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            // Advance only when another position along this axis follows,
686            // and rewind from the last one, so offsets stay reachable.
687            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    // An overflowing element count falls back to the general path, which
815    // reports it as `OffsetOverflow`.
816    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/// Ternary inner loop: `dest[i] = f(a[i], b[i], c[i])`.
999#[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/// Quaternary inner loop: `dest[i] = f(a[i], b[i], c[i], e[i])`.
1046#[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
1109/// Apply a function element-wise from source to destination.
1110///
1111/// The element operation `Op` is applied lazily when reading from `src`.
1112/// Source and destination may have different element types.
1113pub 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    // Small tensor fast path: skip compute_order and compute_block_sizes
1239    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
1300/// Binary element-wise operation: `dest[i] = f(a[i], b[i])`.
1301///
1302/// Source operands `a` and `b` may have different element types from each other
1303/// and from `dest`. The closure `f` handles per-element type conversion.
1304pub 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/// Runtime comparison selected once before entering the element loop.
1357#[non_exhaustive]
1358#[derive(Clone, Copy, Debug, Eq, PartialEq)]
1359pub enum CompareOp {
1360    Eq,
1361    Lt,
1362    Le,
1363    Gt,
1364    Ge,
1365}
1366
1367/// Compare two ordered views elementwise into a Boolean destination.
1368///
1369/// Unlike embedding a runtime operation match inside a [`zip_map2_into`]
1370/// closure, this entry point selects a fixed comparison before traversal.
1371///
1372/// # Errors
1373///
1374/// Returns [`StridedError::ShapeMismatch`] when the source and destination
1375/// shapes differ, or [`StridedError::NonInjectiveOutputLayout`] when distinct
1376/// logical destination elements may overlap.
1377pub 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
1397/// Compare two views into a fully overwritten uninitialized Boolean output.
1398///
1399/// Dtype-independent shape, destination-injectivity, and reachable-byte
1400/// overlap validation completes before the first write. Safe Rust borrows
1401/// already prevent input/output aliasing; the explicit overlap check preserves
1402/// the contract for views produced through unsafe constructors.
1403///
1404/// `Ok(())` means every logical destination element is initialized. An error
1405/// occurs before writes. A panic during replay may leave a partially initialized
1406/// destination, which remains safe to drop as `MaybeUninit<bool>`.
1407///
1408/// # Errors
1409///
1410/// Returns [`StridedError::ShapeMismatch`] for unequal shapes,
1411/// [`StridedError::NonInjectiveOutputLayout`] for an overlapping output layout,
1412/// [`StridedError::OverlappingInputOutput`] for aliased storage, or
1413/// [`StridedError::OffsetOverflow`] when a reachable byte range is not
1414/// representable.
1415pub 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    // Small tensor fast path: skip compute_order and compute_block_sizes
1571    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
1758/// Element-wise multiplication: `dest[i] = a[i] * b[i]`.
1759///
1760/// All views must have the same shape. Broadcast operands should be represented
1761/// as stride-0 views before calling this function.
1762pub 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
1793/// Multiply two views into a fully overwritten uninitialized output.
1794///
1795/// Shape, destination-injectivity, and reachable-byte overlap validation
1796/// completes before the first write. Safe Rust borrows already prevent
1797/// input/output aliasing; the explicit overlap check preserves the contract for
1798/// views produced through unsafe constructors.
1799///
1800/// `Ok(())` means every logical destination element is initialized. An error
1801/// occurs before writes. A panic during replay may leave a partially initialized
1802/// destination, which remains safe to drop as `MaybeUninit<D>`.
1803///
1804/// # Errors
1805///
1806/// Returns [`StridedError::ShapeMismatch`] for unequal shapes,
1807/// [`StridedError::NonInjectiveOutputLayout`] for an overlapping output layout,
1808/// [`StridedError::OverlappingInputOutput`] for aliased storage, or
1809/// [`StridedError::OffsetOverflow`] when a reachable byte range is not
1810/// representable.
1811pub 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
1904/// Broadcasted element-wise multiplication: `dest[i] = a[i] * b[i]`.
1905///
1906/// `a_axes` and `b_axes` map each source axis to an axis of `dest`. Output axes
1907/// not referenced by a source operand are treated as stride-0 broadcast axes.
1908pub 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
1943/// Broadcast and multiply into a fully overwritten uninitialized output.
1944///
1945/// Axis mappings, shapes, destination injectivity, and reachable-byte overlap
1946/// are validated before the first write. Safe Rust borrows already prevent
1947/// input/output aliasing; the explicit overlap check preserves the contract for
1948/// views produced through unsafe constructors.
1949///
1950/// `Ok(())` means every logical destination element is initialized. An error
1951/// occurs before writes. A panic during replay may leave a partially initialized
1952/// destination, which remains safe to drop as `MaybeUninit<D>`.
1953///
1954/// # Errors
1955///
1956/// Returns a typed rank, axis, or shape error for an invalid broadcast mapping,
1957/// [`StridedError::NonInjectiveOutputLayout`] for an overlapping output layout,
1958/// [`StridedError::OverlappingInputOutput`] for aliased storage, or
1959/// [`StridedError::OffsetOverflow`] when a reachable byte range is not
1960/// representable.
1961pub 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
2004/// Ternary element-wise operation: `dest[i] = f(a[i], b[i], c[i])`.
2005pub 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    // Small tensor fast path: skip compute_order and compute_block_sizes
2081    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
2155/// Quaternary element-wise operation: `dest[i] = f(a[i], b[i], c[i], e[i])`.
2156pub 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    // Small tensor fast path: skip compute_order and compute_block_sizes
2260    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;