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