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