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