Skip to main content

strided_basic/
erased.rs

1use crate::erased_common::*;
2use crate::*;
3use core::mem::MaybeUninit;
4use num_complex::{Complex32, Complex64};
5use num_traits::{One, Zero};
6mod arg_reduce;
7mod line;
8mod norm;
9mod reduce_kernel;
10mod scan;
11
12pub use arg_reduce::{ArgReduceOp, ErasedArgReducePlan};
13pub use norm::{ErasedNormPlan, NormKind, NormSpec};
14pub use scan::{ErasedScanPlan, ScanOp, ScanOptions};
15
16use reduce_kernel::{
17    reduce_axes_range, reduce_full_range, AxesRange, FullTraversal, MaxKernel, MinKernel,
18    ProductKernel, ReduceKernel, SumKernel, SumSquaresKernel,
19};
20
21trait ReduceWriter<T> {
22    fn offset(&self) -> isize;
23    /// # Safety
24    /// The pointer may only be used within the validated destination extent.
25    unsafe fn ptr(&mut self) -> *mut T;
26    fn extent(&self) -> usize;
27    /// # Safety
28    /// The offset must be an in-bounds logical reduction destination offset.
29    unsafe fn write_at(&mut self, offset: isize, value: T) {
30        debug_assert!(offset >= 0 && (offset as usize) < self.extent());
31        // SAFETY: reduction layout validation proves the logical offset.
32        unsafe { self.ptr().offset(offset).write(value) }
33    }
34}
35
36struct RawReduceWriter<'a, T> {
37    ptr: *mut T,
38    extent: usize,
39    offset: isize,
40    _marker: core::marker::PhantomData<&'a mut [MaybeUninit<T>]>,
41}
42
43impl<'a, T> ReduceWriter<T> for RawReduceWriter<'a, T> {
44    fn offset(&self) -> isize {
45        self.offset
46    }
47    unsafe fn ptr(&mut self) -> *mut T {
48        self.ptr
49    }
50    fn extent(&self) -> usize {
51        self.extent
52    }
53}
54
55/// Dtype-erased wrapper around [`CopyPlan`].
56#[derive(Clone, Debug)]
57pub struct ErasedCopyPlan {
58    dtype: KernelDType,
59    plan: CopyPlan,
60}
61
62/// Dtype-erased concatenate wrapper.
63#[derive(Clone, Debug)]
64pub struct ErasedConcatenatePlan {
65    dtype: KernelDType,
66    plan: ConcatenatePlan,
67}
68
69impl ErasedCopyPlan {
70    /// Compile a copy plan for one dtype and layout pair.
71    pub fn compile(
72        dtype: KernelDType,
73        dims: &[usize],
74        dst_strides: &[isize],
75        src_strides: &[isize],
76    ) -> Result<Self> {
77        Ok(Self {
78            dtype,
79            plan: CopyPlan::compile(dims, dst_strides, src_strides)?,
80        })
81    }
82
83    #[inline]
84    pub fn dtype(&self) -> KernelDType {
85        self.dtype
86    }
87
88    /// `dest = src` through a non-generic dtype-erased replay boundary.
89    pub fn execute(
90        &self,
91        ctx: &ExecContext,
92        dest: &mut ErasedRawStridedMut<'_>,
93        src: &ErasedRawStridedRef<'_>,
94    ) -> Result<()> {
95        self.check_dtype(dest.dtype())?;
96        self.check_dtype(src.dtype())?;
97
98        let result = ctx.run(|| match self.dtype {
99            KernelDType::F32 => execute_copy::<f32>(&self.plan, dest, src),
100            KernelDType::F64 => execute_copy::<f64>(&self.plan, dest, src),
101            KernelDType::I32 => execute_copy::<i32>(&self.plan, dest, src),
102            KernelDType::I64 => execute_copy::<i64>(&self.plan, dest, src),
103            KernelDType::Bool => execute_copy::<bool>(&self.plan, dest, src),
104            KernelDType::C32 => execute_copy::<Complex32>(&self.plan, dest, src),
105            KernelDType::C64 => execute_copy::<Complex64>(&self.plan, dest, src),
106            _ => Err(StridedError::UnsupportedDType {
107                dtype: self.dtype.label(),
108            }),
109        });
110        result
111    }
112
113    fn check_dtype(&self, actual: KernelDType) -> Result<()> {
114        if actual != self.dtype {
115            return Err(StridedError::DTypeMismatch {
116                expected: self.dtype.label(),
117                actual: actual.label(),
118            });
119        }
120        Ok(())
121    }
122}
123
124impl ErasedConcatenatePlan {
125    /// Validate and store a concatenate plan for one dtype and fixed layout set.
126    pub fn compile(
127        dtype: KernelDType,
128        input_dims: &[&[usize]],
129        input_strides: &[&[isize]],
130        dest_dims: &[usize],
131        dest_strides: &[isize],
132        axis: usize,
133    ) -> Result<Self> {
134        check_static_indexing_dtype(dtype)?;
135        Ok(Self {
136            dtype,
137            plan: ConcatenatePlan::compile(
138                input_dims,
139                input_strides,
140                dest_dims,
141                dest_strides,
142                axis,
143            )?,
144        })
145    }
146
147    #[inline]
148    pub fn dtype(&self) -> KernelDType {
149        self.dtype
150    }
151
152    #[inline]
153    pub fn plan(&self) -> &ConcatenatePlan {
154        &self.plan
155    }
156
157    /// Execute concatenate into an erased output descriptor.
158    pub fn execute(
159        &self,
160        ctx: &ExecContext,
161        dest: &mut ErasedRawStridedMut<'_>,
162        inputs: &[ErasedRawStridedRef<'_>],
163    ) -> Result<()> {
164        check_dtype(self.dtype, dest.dtype())?;
165        for input in inputs {
166            check_dtype(self.dtype, input.dtype())?;
167        }
168
169        let result = ctx.run(|| match self.dtype {
170            KernelDType::F32 => execute_concatenate::<f32>(&self.plan, dest, inputs),
171            KernelDType::F64 => execute_concatenate::<f64>(&self.plan, dest, inputs),
172            KernelDType::I32 => execute_concatenate::<i32>(&self.plan, dest, inputs),
173            KernelDType::I64 => execute_concatenate::<i64>(&self.plan, dest, inputs),
174            KernelDType::Bool => execute_concatenate::<bool>(&self.plan, dest, inputs),
175            KernelDType::C32 => execute_concatenate::<Complex32>(&self.plan, dest, inputs),
176            KernelDType::C64 => execute_concatenate::<Complex64>(&self.plan, dest, inputs),
177            _ => Err(StridedError::UnsupportedDType {
178                dtype: self.dtype.label(),
179            }),
180        });
181        result
182    }
183
184    /// Execute concatenate as a full overwrite of uninitialized output storage.
185    /// On success, every reachable destination slot is fully overwritten;
186    /// unreachable holes are neither read nor initialized. Validation errors
187    /// are returned before any destination write. A panic during execution
188    /// may leave a partially initialized `MaybeUninit` destination, which is
189    /// still safely droppable; no readable value is promised for unwritten
190    /// reachable slots.
191    pub fn execute_uninit(
192        &self,
193        ctx: &ExecContext,
194        dest: &mut ErasedRawStridedUninitMut<'_>,
195        inputs: &[ErasedRawStridedPtr<'_>],
196    ) -> Result<()> {
197        check_dtype(self.dtype, dest.dtype())?;
198        if inputs.len() != self.plan.input_count() {
199            return Err(StridedError::RankMismatch(
200                inputs.len(),
201                self.plan.input_count(),
202            ));
203        }
204        for input in inputs {
205            check_dtype(self.dtype, input.dtype())?;
206        }
207        for (position, input) in inputs.iter().enumerate() {
208            validate_uninit_no_overlap(dest, input, position)?;
209        }
210        for input in inputs {
211            // SAFETY: input/output overlap was rejected before forming references.
212            unsafe { input.try_as_ref_after_no_overlap() }?;
213        }
214
215        ctx.run(|| match self.dtype {
216            KernelDType::F32 => execute_concatenate_uninit::<f32>(&self.plan, dest, inputs),
217            KernelDType::F64 => execute_concatenate_uninit::<f64>(&self.plan, dest, inputs),
218            KernelDType::I32 => execute_concatenate_uninit::<i32>(&self.plan, dest, inputs),
219            KernelDType::I64 => execute_concatenate_uninit::<i64>(&self.plan, dest, inputs),
220            KernelDType::Bool => execute_concatenate_uninit::<bool>(&self.plan, dest, inputs),
221            KernelDType::C32 => execute_concatenate_uninit::<Complex32>(&self.plan, dest, inputs),
222            KernelDType::C64 => execute_concatenate_uninit::<Complex64>(&self.plan, dest, inputs),
223            _ => Err(StridedError::UnsupportedDType {
224                dtype: self.dtype.label(),
225            }),
226        })
227    }
228}
229
230/// Runtime reduction operation for dtype-erased full reductions.
231#[non_exhaustive]
232#[derive(Clone, Copy, Debug, Eq, PartialEq)]
233pub enum ReduceOp {
234    Sum,
235    Product,
236    /// Sum of same-dtype rounded squares.
237    ///
238    /// Each input is first multiplied by itself without FMA contraction, then
239    /// accumulated under the same association policy as [`Self::Sum`].
240    SumSquares,
241    /// NaN-propagating maximum for `f32`/`f64`, ordered maximum for
242    /// `i32`/`i64`.
243    ///
244    /// Any NaN operand makes the result the canonical NaN of the dtype. The
245    /// identity written for an empty reduction is `-inf` for floats and the
246    /// dtype minimum for integers. Complex and `bool` dtypes are rejected.
247    /// When the reduced values contain both `-0.0` and `+0.0`, which zero is
248    /// returned is unspecified.
249    Max,
250    /// NaN-propagating minimum for `f32`/`f64`, ordered minimum for
251    /// `i32`/`i64`.
252    ///
253    /// Any NaN operand makes the result the canonical NaN of the dtype. The
254    /// identity written for an empty reduction is `+inf` for floats and the
255    /// dtype maximum for integers. Complex and `bool` dtypes are rejected.
256    /// When the reduced values contain both `-0.0` and `+0.0`, which zero is
257    /// returned is unspecified.
258    Min,
259}
260
261/// Dtype-erased reduction wrapper.
262///
263/// This is the erased replay boundary for full-tensor scalar reductions and
264/// axis reductions with a fixed output layout. It supports only operations with
265/// an unambiguous identity value in the selected dtype.
266///
267/// # Evaluation order
268///
269/// The op is matched once per execution; every loop runs a kernel
270/// specialized for one `(dtype, op)` pair. Integer results are exact
271/// (wrapping) and [`ReduceOp::Max`] / [`ReduceOp::Min`] do not depend on
272/// order, so the rules below only affect the rounding of floating point and
273/// complex [`ReduceOp::Sum`], [`ReduceOp::Product`] and
274/// [`ReduceOp::SumSquares`]. Two details of float max and min follow the
275/// visit order: a NaN input makes the result NaN, but which NaN (payload and
276/// sign) is returned is unspecified, and when the extremum is a zero its sign
277/// is unspecified if both `+0.0` and `-0.0` occur.
278///
279/// * Full reductions (from [`Self::compile`], or [`Self::compile_axes`]
280///   with every axis reduced) visit the source in a layout derived order:
281///   extent one axes are dropped, the rest are ordered by increasing
282///   absolute stride and contiguous axes are fused into runs. Each run is
283///   reduced by the SIMD kernel (unit stride, `simd` feature), a sixteen
284///   lane kernel (other unit stride runs) or an eight lane kernel (strided
285///   runs), and runs are combined left to right. A dense column major
286///   source therefore reduces exactly like one contiguous slice.
287/// * With a parallel context and more than `MINTHREADLENGTH` elements, a
288///   full reduction splits the traversal into one deterministic range per
289///   worker thread. Each range runs the same run kernel as the serial path
290///   and the partials are combined in a fixed tree. The result is bitwise
291///   repeatable for a fixed layout and thread count, but may differ between
292///   thread counts and from the serial result.
293/// * Axis reductions with kept axes compute every output independently, so
294///   the result never depends on the context or thread count. Each output
295///   folds its reduced elements sequentially in caller reduced axis order,
296///   except that a unit stride leading reduced run (after fusing adjacent
297///   reduced axes) of length at least 16 is reduced by the contiguous run
298///   kernel, with runs folded in order. When instead the leading kept axis
299///   has unit stride, adjacent outputs are accumulated as a block from one
300///   contiguous source run per reduced element; each output still folds in
301///   caller reduced axis order.
302#[derive(Clone, Debug)]
303pub struct ErasedReducePlan {
304    dtype: KernelDType,
305    op: ReduceOp,
306    layout: ReduceLayout,
307}
308
309#[derive(Clone, Debug)]
310enum ReduceLayout {
311    Full {
312        dims: Vec<usize>,
313        src_strides: Vec<isize>,
314        traversal: FullTraversal,
315    },
316    Axes {
317        src_dims: Vec<usize>,
318        src_strides: Vec<isize>,
319        dest_dims: Vec<usize>,
320        dest_strides: Vec<isize>,
321        outer_axes: Vec<ReduceOuterAxis>,
322        inner_axes: Vec<ReduceInnerAxis>,
323        dest_total: usize,
324        reduce_total: usize,
325        /// Set when every source axis is reduced; the plan then runs the
326        /// full-reduction traversal.
327        full: Option<FullTraversal>,
328    },
329}
330
331#[derive(Clone, Copy, Debug)]
332struct ReduceOuterAxis {
333    extent: usize,
334    source_step: isize,
335    source_reset: isize,
336    dest_step: isize,
337    dest_reset: isize,
338}
339#[derive(Clone, Copy, Debug)]
340struct ReduceInnerAxis {
341    extent: usize,
342    source_step: isize,
343    source_reset: isize,
344}
345impl ReduceLayout {
346    fn src_dims(&self) -> &[usize] {
347        match self {
348            Self::Full { dims, .. } => dims,
349            Self::Axes { src_dims, .. } => src_dims,
350        }
351    }
352
353    fn src_strides(&self) -> &[isize] {
354        match self {
355            Self::Full { src_strides, .. } | Self::Axes { src_strides, .. } => src_strides,
356        }
357    }
358
359    fn check_src_layout(&self, src: &ErasedRawStridedRef<'_>) -> Result<()> {
360        if src.dims() != self.src_dims() || src.strides() != self.src_strides() {
361            return Err(StridedError::PlanLayoutMismatch);
362        }
363        Ok(())
364    }
365}
366
367#[derive(Clone, Copy, Debug)]
368struct AxesLayout<'a> {
369    outer_axes: &'a [ReduceOuterAxis],
370    inner_axes: &'a [ReduceInnerAxis],
371    dest_total: usize,
372    reduce_total: usize,
373}
374
375impl ErasedReducePlan {
376    /// Validate and store a full-reduction plan for one dtype and source layout.
377    pub fn compile(
378        dtype: KernelDType,
379        op: ReduceOp,
380        dims: &[usize],
381        src_strides: &[isize],
382    ) -> Result<Self> {
383        check_reduce_op_dtype(dtype, op)?;
384        if dims.len() != src_strides.len() {
385            return Err(StridedError::StrideLengthMismatch);
386        }
387        checked_total_len(dims)?;
388        let traversal = FullTraversal::compile(dims, src_strides)?;
389        Ok(Self {
390            dtype,
391            op,
392            layout: ReduceLayout::Full {
393                dims: dims.to_vec(),
394                src_strides: src_strides.to_vec(),
395                traversal,
396            },
397        })
398    }
399
400    /// Validate and store an axis-reduction plan for one dtype and fixed source/output layouts.
401    ///
402    /// `axes` names the source axes reduced away. Output dimensions must be the
403    /// remaining source dimensions in source-axis order. When all axes are
404    /// reduced, any output layout with exactly one reachable element is accepted.
405    #[allow(clippy::too_many_arguments)]
406    pub fn compile_axes(
407        dtype: KernelDType,
408        op: ReduceOp,
409        src_dims: &[usize],
410        src_strides: &[isize],
411        dest_dims: &[usize],
412        dest_strides: &[isize],
413        axes: &[usize],
414    ) -> Result<Self> {
415        check_reduce_op_dtype(dtype, op)?;
416        if src_dims.len() != src_strides.len() || dest_dims.len() != dest_strides.len() {
417            return Err(StridedError::StrideLengthMismatch);
418        }
419        checked_total_len(src_dims)?;
420        check_reduce_layout_offset_arithmetic(src_dims, src_strides)?;
421        let dest_total = checked_total_len(dest_dims)?;
422        check_reduce_layout_offset_arithmetic(dest_dims, dest_strides)?;
423        if !crate::layout_check::is_injective_layout(dest_dims, dest_strides) {
424            return Err(StridedError::NonInjectiveOutputLayout);
425        }
426        validate_unique_axes(axes, src_dims.len())?;
427
428        let kept_axes: Vec<usize> = (0..src_dims.len())
429            .filter(|axis| !axes.contains(axis))
430            .collect();
431        let expected_dest_dims: Vec<usize> = kept_axes.iter().map(|&axis| src_dims[axis]).collect();
432        if expected_dest_dims.is_empty() {
433            if dest_total != 1 {
434                return Err(StridedError::ShapeMismatch(
435                    dest_dims.to_vec(),
436                    expected_dest_dims,
437                ));
438            }
439        } else if dest_dims != expected_dest_dims.as_slice() {
440            return Err(StridedError::ShapeMismatch(
441                dest_dims.to_vec(),
442                expected_dest_dims,
443            ));
444        }
445
446        let reduce_total = axes
447            .iter()
448            .try_fold(1usize, |total, &axis| total.checked_mul(src_dims[axis]))
449            .ok_or(StridedError::OffsetOverflow)?;
450        let outer_axes = compress_reduce_outer_axes(
451            kept_axes
452                .iter()
453                .enumerate()
454                .map(|(dest_axis, &src_axis)| {
455                    let extent = src_dims[src_axis];
456                    Ok(ReduceOuterAxis {
457                        extent,
458                        source_step: src_strides[src_axis],
459                        source_reset: checked_reduce_reset(extent, src_strides[src_axis])?,
460                        dest_step: dest_strides[dest_axis],
461                        dest_reset: checked_reduce_reset(extent, dest_strides[dest_axis])?,
462                    })
463                })
464                .collect::<Result<Vec<_>>>()?,
465        )?;
466        let inner_axes = compress_reduce_inner_axes(
467            axes.iter()
468                .map(|&src_axis| {
469                    let extent = src_dims[src_axis];
470                    Ok(ReduceInnerAxis {
471                        extent,
472                        source_step: src_strides[src_axis],
473                        source_reset: checked_reduce_reset(extent, src_strides[src_axis])?,
474                    })
475                })
476                .collect::<Result<Vec<_>>>()?,
477        )?;
478        let full = if kept_axes.is_empty() {
479            Some(FullTraversal::compile(src_dims, src_strides)?)
480        } else {
481            None
482        };
483        Ok(Self {
484            dtype,
485            op,
486            layout: ReduceLayout::Axes {
487                src_dims: src_dims.to_vec(),
488                src_strides: src_strides.to_vec(),
489                dest_dims: dest_dims.to_vec(),
490                dest_strides: dest_strides.to_vec(),
491                outer_axes,
492                inner_axes,
493                dest_total,
494                reduce_total,
495                full,
496            },
497        })
498    }
499
500    #[inline]
501    pub fn dtype(&self) -> KernelDType {
502        self.dtype
503    }
504
505    #[inline]
506    pub fn op(&self) -> ReduceOp {
507        self.op
508    }
509
510    /// Execute the reduction into an erased output descriptor.
511    pub fn execute(
512        &self,
513        ctx: &ExecContext,
514        dest: &mut ErasedRawStridedMut<'_>,
515        src: &ErasedRawStridedRef<'_>,
516    ) -> Result<()> {
517        check_dtype(self.dtype, dest.dtype())?;
518        check_dtype(self.dtype, src.dtype())?;
519        self.layout.check_src_layout(src)?;
520        match &self.layout {
521            ReduceLayout::Full { .. } => {
522                let dest_len = checked_total_len(dest.dims())?;
523                if dest_len != 1 {
524                    return Err(StridedError::RankMismatch(dest_len, 1));
525                }
526            }
527            ReduceLayout::Axes {
528                dest_dims,
529                dest_strides,
530                ..
531            } => {
532                if dest.dims() != dest_dims.as_slice() || dest.strides() != dest_strides.as_slice()
533                {
534                    return Err(StridedError::PlanLayoutMismatch);
535                }
536            }
537        }
538
539        let result = match self.dtype {
540            KernelDType::F32 => {
541                let mut writer = reduce_writer::<f32>(dest)?;
542                dispatch_reduce::<f32, _>(self.op, &self.layout, ctx, &mut writer, src)
543            }
544            KernelDType::F64 => {
545                let mut writer = reduce_writer::<f64>(dest)?;
546                dispatch_reduce::<f64, _>(self.op, &self.layout, ctx, &mut writer, src)
547            }
548            KernelDType::I32 => {
549                let mut writer = reduce_writer::<i32>(dest)?;
550                dispatch_reduce::<i32, _>(self.op, &self.layout, ctx, &mut writer, src)
551            }
552            KernelDType::I64 => {
553                let mut writer = reduce_writer::<i64>(dest)?;
554                dispatch_reduce::<i64, _>(self.op, &self.layout, ctx, &mut writer, src)
555            }
556            KernelDType::C32 => {
557                let mut writer = reduce_writer::<Complex32>(dest)?;
558                dispatch_reduce::<Complex32, _>(self.op, &self.layout, ctx, &mut writer, src)
559            }
560            KernelDType::C64 => {
561                let mut writer = reduce_writer::<Complex64>(dest)?;
562                dispatch_reduce::<Complex64, _>(self.op, &self.layout, ctx, &mut writer, src)
563            }
564            _ => Err(StridedError::UnsupportedDType {
565                dtype: self.dtype.label(),
566            }),
567        };
568        result
569    }
570
571    /// On success, every reachable destination slot is fully overwritten;
572    /// unreachable holes are neither read nor initialized. Validation errors
573    /// are returned before any destination write. A panic during execution
574    /// may leave a partially initialized `MaybeUninit` destination, which is
575    /// still safely droppable; no readable value is promised for unwritten
576    /// reachable slots.
577    pub fn execute_uninit(
578        &self,
579        ctx: &ExecContext,
580        dest: &mut ErasedRawStridedUninitMut<'_>,
581        src: &ErasedRawStridedPtr<'_>,
582    ) -> Result<()> {
583        check_dtype(self.dtype, dest.dtype())?;
584        check_dtype(self.dtype, src.dtype())?;
585        validate_uninit_no_overlap(dest, src, 0)?;
586        // SAFETY: the owning erased entry rejected all input/output overlap before conversion.
587        let src = unsafe { src.try_as_ref_after_no_overlap() }?;
588        self.layout.check_src_layout(&src)?;
589        match &self.layout {
590            ReduceLayout::Full { .. } => {
591                let total = checked_total_len(dest.dims())?;
592                if total != 1 {
593                    return Err(StridedError::RankMismatch(total, 1));
594                }
595            }
596            ReduceLayout::Axes {
597                dest_dims,
598                dest_strides,
599                ..
600            } => {
601                if dest.dims() != dest_dims.as_slice() || dest.strides() != dest_strides.as_slice()
602                {
603                    return Err(StridedError::PlanLayoutMismatch);
604                }
605            }
606        }
607        macro_rules! run {
608            ($ty:ty) => {{
609                let mut writer = reduce_uninit_writer::<$ty>(dest)?;
610                dispatch_reduce::<$ty, _>(self.op, &self.layout, ctx, &mut writer, &src)
611            }};
612        }
613        match self.dtype {
614            KernelDType::F32 => run!(f32),
615            KernelDType::F64 => run!(f64),
616            KernelDType::I32 => run!(i32),
617            KernelDType::I64 => run!(i64),
618            KernelDType::C32 => run!(Complex32),
619            KernelDType::C64 => run!(Complex64),
620            _ => Err(StridedError::UnsupportedDType {
621                dtype: self.dtype.label(),
622            }),
623        }
624    }
625}
626
627fn reduce_writer<'a, T>(dest: &'a mut ErasedRawStridedMut<'_>) -> Result<RawReduceWriter<'a, T>>
628where
629    T: KernelStorageElement,
630{
631    let offset = dest.offset();
632    let data = dest.data_as_mut::<T>()?;
633    let ptr = data.as_mut_ptr();
634    let extent = data.len();
635    Ok(RawReduceWriter {
636        ptr,
637        extent,
638        offset,
639        _marker: core::marker::PhantomData,
640    })
641}
642
643fn reduce_uninit_writer<'a, T>(
644    dest: &'a mut ErasedRawStridedUninitMut<'_>,
645) -> Result<RawReduceWriter<'a, T>>
646where
647    T: KernelStorageElement,
648{
649    let offset = dest.offset();
650    let data = dest.data_as_uninit_mut::<T>()?;
651    let ptr = data.as_mut_ptr().cast::<T>();
652    let extent = data.len();
653    Ok(RawReduceWriter {
654        ptr,
655        extent,
656        offset,
657        _marker: core::marker::PhantomData,
658    })
659}
660
661fn check_reduce_dtype(dtype: KernelDType) -> Result<()> {
662    match dtype {
663        KernelDType::F32
664        | KernelDType::F64
665        | KernelDType::I32
666        | KernelDType::I64
667        | KernelDType::C32
668        | KernelDType::C64 => Ok(()),
669        _ => Err(StridedError::UnsupportedDType {
670            dtype: dtype.label(),
671        }),
672    }
673}
674
675fn check_reduce_op_dtype(dtype: KernelDType, op: ReduceOp) -> Result<()> {
676    if op == ReduceOp::SumSquares && !matches!(dtype, KernelDType::F32 | KernelDType::F64) {
677        return Err(StridedError::UnsupportedDType {
678            dtype: dtype.label(),
679        });
680    }
681    if matches!(op, ReduceOp::Max | ReduceOp::Min)
682        && !matches!(
683            dtype,
684            KernelDType::F32 | KernelDType::F64 | KernelDType::I32 | KernelDType::I64
685        )
686    {
687        return Err(StridedError::UnsupportedDType {
688            dtype: dtype.label(),
689        });
690    }
691    check_reduce_dtype(dtype)
692}
693
694fn checked_total_len(dims: &[usize]) -> Result<usize> {
695    if dims.is_empty() {
696        return Ok(1);
697    }
698    dims.iter()
699        .try_fold(1usize, |acc, &dim| acc.checked_mul(dim))
700        .ok_or(StridedError::OffsetOverflow)
701}
702
703fn execute_copy<T>(
704    plan: &CopyPlan,
705    dest: &mut ErasedRawStridedMut<'_>,
706    src: &ErasedRawStridedRef<'_>,
707) -> Result<()>
708where
709    T: Copy + crate::MaybeSendSync + KernelStorageElement,
710{
711    let source_data = src.data_as::<T>()?;
712    let dest_dims = dest.dims();
713    let dest_strides = dest.strides();
714    let dest_offset = dest.offset();
715    let dest_data = dest.data_as_mut::<T>()?;
716    let source = unsafe {
717        RawStridedRef::new_unchecked(source_data, src.dims(), src.strides(), src.offset())
718    };
719    let mut dest =
720        unsafe { RawStridedMut::new_unchecked(dest_data, dest_dims, dest_strides, dest_offset) };
721    plan.execute(&mut dest, &source)
722}
723
724fn execute_concatenate<T>(
725    plan: &ConcatenatePlan,
726    dest: &mut ErasedRawStridedMut<'_>,
727    inputs: &[ErasedRawStridedRef<'_>],
728) -> Result<()>
729where
730    T: Copy + crate::MaybeSendSync + KernelStorageElement,
731{
732    if inputs.len() != plan.input_count() {
733        return Err(StridedError::RankMismatch(inputs.len(), plan.input_count()));
734    }
735    let dest_dims = dest.dims();
736    let dest_strides = dest.strides();
737    let dest_offset = dest.offset();
738    let dest_data = dest.data_as_mut::<T>()?;
739    let mut dest_ref =
740        unsafe { RawStridedMut::new_unchecked(dest_data, dest_dims, dest_strides, dest_offset) };
741    plan.check_dest_layout(&dest_ref)?;
742
743    for (position, input) in inputs.iter().enumerate() {
744        let input_data = input.data_as::<T>()?;
745        let input_ref = unsafe {
746            RawStridedRef::new_unchecked(input_data, input.dims(), input.strides(), input.offset())
747        };
748        plan.check_input_layout(position, &input_ref)?;
749        plan.segment_offset(position, dest_offset)?;
750    }
751    if plan.prefers_whole_plan() {
752        let input_refs = inputs
753            .iter()
754            .map(|input| {
755                let input_data = input.data_as::<T>()?;
756                Ok(unsafe {
757                    RawStridedRef::new_unchecked(
758                        input_data,
759                        input.dims(),
760                        input.strides(),
761                        input.offset(),
762                    )
763                })
764            })
765            .collect::<Result<Vec<_>>>()?;
766        // Splits the concatenated index space across and within segments.
767        return plan.execute(&mut dest_ref, &input_refs);
768    }
769    for (position, input) in inputs.iter().enumerate() {
770        let input_data = input.data_as::<T>()?;
771        let input_ref = unsafe {
772            RawStridedRef::new_unchecked(input_data, input.dims(), input.strides(), input.offset())
773        };
774        plan.execute_segment(position, &mut dest_ref, &input_ref)?;
775    }
776    Ok(())
777}
778
779fn execute_concatenate_uninit<T>(
780    plan: &ConcatenatePlan,
781    dest: &mut ErasedRawStridedUninitMut<'_>,
782    inputs: &[ErasedRawStridedPtr<'_>],
783) -> Result<()>
784where
785    T: Copy + crate::MaybeSendSync + KernelStorageElement,
786{
787    let dest_dims = dest.dims();
788    let dest_strides = dest.strides();
789    let dest_offset = dest.offset();
790    let dest_data = dest.data_as_uninit_mut::<T>()?;
791    let mut dest_ref =
792        unsafe { RawStridedMut::new_unchecked(dest_data, dest_dims, dest_strides, dest_offset) };
793    plan.check_dest_layout(&dest_ref)?;
794
795    for (position, input) in inputs.iter().enumerate() {
796        // SAFETY: input/output overlap was rejected before forming references.
797        let input = unsafe { input.try_as_ref_after_no_overlap() }?;
798        let input_data = input.data_as::<T>()?;
799        let input_ref = unsafe {
800            RawStridedRef::new_unchecked(input_data, input.dims(), input.strides(), input.offset())
801        };
802        plan.check_input_layout(position, &input_ref)?;
803        plan.segment_offset(position, dest_offset)?;
804    }
805    if plan.prefers_whole_plan() {
806        let inputs = inputs
807            .iter()
808            // SAFETY: input/output overlap was rejected before forming references.
809            .map(|input| unsafe { input.try_as_ref_after_no_overlap() })
810            .collect::<Result<Vec<_>>>()?;
811        let input_refs = inputs
812            .iter()
813            .map(|input| {
814                let input_data = input.data_as::<T>()?;
815                Ok(unsafe {
816                    RawStridedRef::new_unchecked(
817                        input_data,
818                        input.dims(),
819                        input.strides(),
820                        input.offset(),
821                    )
822                })
823            })
824            .collect::<Result<Vec<_>>>()?;
825        // Splits the concatenated index space across and within segments.
826        return plan.execute_uninit(&mut dest_ref, &input_refs);
827    }
828    for (position, input) in inputs.iter().enumerate() {
829        // SAFETY: input/output overlap was rejected before forming references.
830        let input = unsafe { input.try_as_ref_after_no_overlap() }?;
831        let input_data = input.data_as::<T>()?;
832        let input_ref = unsafe {
833            RawStridedRef::new_unchecked(input_data, input.dims(), input.strides(), input.offset())
834        };
835        plan.execute_segment_uninit(position, &mut dest_ref, &input_ref)?;
836    }
837    Ok(())
838}
839
840/// Whether a reduction context runs on the caller thread without entering
841/// the scheduler.
842fn reduce_context_is_serial(ctx: &ExecContext) -> bool {
843    ctx.is_serial()
844        || ctx
845            .max_threads_limit()
846            .is_some_and(|max_threads| max_threads.get() == 1)
847}
848
849fn execute_reduce_full<T, W, K>(
850    ctx: &ExecContext,
851    dest: &mut W,
852    src: &ErasedRawStridedRef<'_>,
853    traversal: &FullTraversal,
854) -> Result<()>
855where
856    T: ErasedReduceScalar,
857    W: ReduceWriter<T>,
858    K: ReduceKernel<T>,
859{
860    let source = src.data_as::<T>()?;
861    let total = traversal.total();
862    let value = if total == 0 {
863        K::identity()
864    } else if reduce_context_is_serial(ctx) {
865        // INVARIANT: (1) FullTraversal::compile checked every step/reset of
866        // the source layout; (2) the raw source descriptor validated every
867        // reachable offset; (3) execute checked exact plan-layout equality.
868        // SAFETY: the three-link invariant proves every traversal offset is
869        // readable, and `0..total` is non-empty.
870        unsafe { reduce_full_range::<T, K>(source.as_ptr(), src.offset(), traversal, 0..total)? }
871    } else {
872        ctx.run(|| reduce_full_policy::<T, K>(source, src.offset(), traversal))?
873    };
874
875    // SAFETY: validated rank-zero destination layout proves the offset.
876    unsafe { dest.write_at(dest.offset(), value) };
877    Ok(())
878}
879
880fn reduce_full_policy<T, K>(
881    source: &[T],
882    source_base: isize,
883    traversal: &FullTraversal,
884) -> Result<T>
885where
886    T: ErasedReduceScalar,
887    K: ReduceKernel<T>,
888{
889    let total = traversal.total();
890    #[cfg(feature = "parallel")]
891    {
892        let nthreads = crate::threading::parallel_threads_for_len(total);
893        if nthreads > 1 {
894            let source_ptr = crate::threading::SendPtr(source.as_ptr() as *mut T);
895            return crate::threading::parallel_map_reduce(
896                0..total,
897                nthreads,
898                &|range| {
899                    // SAFETY: same three-link invariant as the serial call in
900                    // `execute_reduce_full`; each worker range is non-empty.
901                    unsafe {
902                        reduce_full_range::<T, K>(
903                            source_ptr.as_const(),
904                            source_base,
905                            traversal,
906                            range,
907                        )
908                    }
909                },
910                &|left, right| Ok(K::combine(left?, right?)),
911            );
912        }
913    }
914    // SAFETY: same three-link invariant as the serial call in
915    // `execute_reduce_full`; `0..total` is non-empty.
916    unsafe { reduce_full_range::<T, K>(source.as_ptr(), source_base, traversal, 0..total) }
917}
918
919fn dispatch_reduce<T, W>(
920    op: ReduceOp,
921    layout: &ReduceLayout,
922    ctx: &ExecContext,
923    dest: &mut W,
924    src: &ErasedRawStridedRef<'_>,
925) -> Result<()>
926where
927    T: ErasedReduceScalar,
928    W: ReduceWriter<T>,
929{
930    // The op is matched once per execution; every loop below is generic over
931    // the selected kernel type.
932    match op {
933        ReduceOp::Sum => dispatch_reduce_kernel::<T, W, SumKernel>(layout, ctx, dest, src),
934        ReduceOp::Product => dispatch_reduce_kernel::<T, W, ProductKernel>(layout, ctx, dest, src),
935        ReduceOp::SumSquares => {
936            dispatch_reduce_kernel::<T, W, SumSquaresKernel>(layout, ctx, dest, src)
937        }
938        ReduceOp::Max => dispatch_reduce_kernel::<T, W, MaxKernel>(layout, ctx, dest, src),
939        ReduceOp::Min => dispatch_reduce_kernel::<T, W, MinKernel>(layout, ctx, dest, src),
940    }
941}
942
943fn dispatch_reduce_kernel<T, W, K>(
944    layout: &ReduceLayout,
945    ctx: &ExecContext,
946    dest: &mut W,
947    src: &ErasedRawStridedRef<'_>,
948) -> Result<()>
949where
950    T: ErasedReduceScalar,
951    W: ReduceWriter<T>,
952    K: ReduceKernel<T>,
953{
954    match layout {
955        ReduceLayout::Full { traversal, .. } => {
956            execute_reduce_full::<T, W, K>(ctx, dest, src, traversal)
957        }
958        ReduceLayout::Axes {
959            outer_axes,
960            inner_axes,
961            dest_total,
962            reduce_total,
963            full,
964            ..
965        } => {
966            if let Some(traversal) = full {
967                return execute_reduce_full::<T, W, K>(ctx, dest, src, traversal);
968            }
969            execute_reduce_axes::<T, W, K>(
970                ctx,
971                dest,
972                src,
973                AxesLayout {
974                    outer_axes,
975                    inner_axes,
976                    dest_total: *dest_total,
977                    reduce_total: *reduce_total,
978                },
979            )
980        }
981    }
982}
983
984fn execute_reduce_axes<T, W, K>(
985    ctx: &ExecContext,
986    dest: &mut W,
987    src: &ErasedRawStridedRef<'_>,
988    layout: AxesLayout<'_>,
989) -> Result<()>
990where
991    T: ErasedReduceScalar,
992    W: ReduceWriter<T>,
993    K: ReduceKernel<T>,
994{
995    if layout.dest_total == 0 {
996        return Ok(());
997    }
998
999    if layout.reduce_total == 0 {
1000        if ctx.is_serial() {
1001            execute_reduce_axes_identity_serial::<T, W, K>(dest, layout)
1002        } else {
1003            ctx.run(|| execute_reduce_axes_identity_policy::<T, W, K>(dest, layout))
1004        }
1005    } else if reduce_context_is_serial(ctx) {
1006        execute_reduce_axes_policy::<T, W, K>(dest, src, layout, false)
1007    } else {
1008        ctx.run(|| execute_reduce_axes_policy::<T, W, K>(dest, src, layout, true))
1009    }
1010}
1011
1012fn execute_reduce_axes_policy<T, W, K>(
1013    dest: &mut W,
1014    src: &ErasedRawStridedRef<'_>,
1015    layout: AxesLayout<'_>,
1016    allow_parallel: bool,
1017) -> Result<()>
1018where
1019    T: ErasedReduceScalar,
1020    W: ReduceWriter<T>,
1021    K: ReduceKernel<T>,
1022{
1023    let source_data = src.data_as::<T>()?;
1024    let source_base = src.offset();
1025    let dest_base = dest.offset();
1026    // SAFETY: the validated reduction writer owns the destination allocation.
1027    let dest_ptr = unsafe { dest.ptr() };
1028    #[cfg(feature = "parallel")]
1029    if allow_parallel {
1030        // The threshold counts source elements read (outputs times reduced
1031        // elements), as documented in docs/design/erased-execution-policy.md,
1032        // so a few thousand outputs with long reductions still split. Worker
1033        // ranges partition the independent outputs; with the output block
1034        // kernel this splits along the contiguous kept axis.
1035        let work = layout.dest_total.saturating_mul(layout.reduce_total);
1036        let nthreads = crate::threading::parallel_threads_for_len(work).min(layout.dest_total);
1037        if nthreads > 1 {
1038            let dest_ptr = crate::threading::SendPtr(dest_ptr);
1039            let source_ptr = crate::threading::SendPtr(source_data.as_ptr() as *mut T);
1040            return crate::threading::parallel_map_reduce(
1041                0..layout.dest_total,
1042                nthreads,
1043                &|range| {
1044                    let parts = AxesRange {
1045                        source: source_ptr.as_const(),
1046                        source_base,
1047                        dest: dest_ptr.as_ptr(),
1048                        dest_base,
1049                        outer_axes: layout.outer_axes,
1050                        inner_axes: layout.inner_axes,
1051                        reduce_total: layout.reduce_total,
1052                    };
1053                    // SAFETY: see `execute_reduce_axes_policy` serial call.
1054                    unsafe { reduce_axes_range::<T, K>(parts, range) }
1055                },
1056                &|left, right| left.and(right),
1057            );
1058        }
1059    }
1060    #[cfg(not(feature = "parallel"))]
1061    let _ = allow_parallel;
1062
1063    let parts = AxesRange {
1064        source: source_data.as_ptr(),
1065        source_base,
1066        dest: dest_ptr,
1067        dest_base,
1068        outer_axes: layout.outer_axes,
1069        inner_axes: layout.inner_axes,
1070        reduce_total: layout.reduce_total,
1071    };
1072    // INVARIANT: (1) compile_axes checked signed source/destination spans and
1073    // every cursor step/reset; (2) raw input and output descriptors validated
1074    // every reachable offset; (3) execute checked exact plan-layout equality
1075    // before dispatch. Output ranges are disjoint, and the initialized entry
1076    // holds distinct borrows while the uninit entry rejected overlap.
1077    // SAFETY: the three-link invariant above; `reduce_total` is nonzero.
1078    unsafe { reduce_axes_range::<T, K>(parts, 0..layout.dest_total) }
1079}
1080
1081fn execute_reduce_axes_identity_serial<T, W, K>(dest: &mut W, layout: AxesLayout<'_>) -> Result<()>
1082where
1083    T: ErasedReduceScalar,
1084    W: ReduceWriter<T>,
1085    K: ReduceKernel<T>,
1086{
1087    let mut outer = ReduceOuterCursor::decode(0, 0, dest.offset(), layout.outer_axes)?;
1088    for output in 0..layout.dest_total {
1089        // INVARIANT: (1) compile_axes checked the destination span and every
1090        // destination step/reset, including -(extent-1)*stride; (2) the raw
1091        // destination descriptor validated every reachable offset; (3) execute
1092        // checked exact plan-layout equality before dispatch.
1093        // SAFETY: the three-link layout invariant proves this destination
1094        // cursor offset is in bounds; the source pointer is intentionally never
1095        // formed for an empty reduction domain.
1096        unsafe { dest.write_at(outer.dest_offset, K::identity()) };
1097        if output + 1 < layout.dest_total {
1098            outer.advance();
1099        }
1100    }
1101    Ok(())
1102}
1103fn execute_reduce_axes_identity_policy<T, W, K>(dest: &mut W, layout: AxesLayout<'_>) -> Result<()>
1104where
1105    T: ErasedReduceScalar,
1106    W: ReduceWriter<T>,
1107    K: ReduceKernel<T>,
1108{
1109    #[cfg(feature = "parallel")]
1110    {
1111        let nthreads = crate::threading::parallel_threads_for_len(layout.dest_total);
1112        if nthreads > 1 {
1113            return execute_reduce_axes_identity_parallel::<T, W, K>(dest, layout, nthreads);
1114        }
1115    }
1116    execute_reduce_axes_identity_serial::<T, W, K>(dest, layout)
1117}
1118#[cfg(feature = "parallel")]
1119fn execute_reduce_axes_identity_parallel<T, W, K>(
1120    dest: &mut W,
1121    layout: AxesLayout<'_>,
1122    nthreads: usize,
1123) -> Result<()>
1124where
1125    T: ErasedReduceScalar,
1126    W: ReduceWriter<T>,
1127    K: ReduceKernel<T>,
1128{
1129    // SAFETY: the validated reduction writer owns the destination allocation.
1130    let dest_ptr = crate::threading::SendPtr(unsafe { dest.ptr() });
1131    let dest_offset_base = dest.offset();
1132    crate::threading::parallel_map_reduce(
1133        0..layout.dest_total,
1134        nthreads,
1135        &|range| {
1136            let range_end = range.end;
1137            let mut outer =
1138                ReduceOuterCursor::decode(range.start, 0, dest_offset_base, layout.outer_axes)?;
1139            let dest_ptr = dest_ptr.as_ptr();
1140            for output in range {
1141                // INVARIANT: (1) compile_axes checked the destination span and
1142                // every destination step/reset, including
1143                // -(extent-1)*stride; (2) the raw destination descriptor
1144                // validated every reachable pointer offset; (3) execute checked
1145                // exact plan-layout equality before dispatch.
1146                // SAFETY: the three-link layout invariant proves this
1147                // destination offset is in bounds; no source pointer is formed.
1148                unsafe { dest_ptr.offset(outer.dest_offset).write(K::identity()) };
1149                if output + 1 < range_end {
1150                    outer.advance();
1151                }
1152            }
1153            Ok(())
1154        },
1155        &|left, right| left.and(right),
1156    )
1157}
1158
1159trait ErasedReduceScalar:
1160    KernelStorageElement
1161    + Copy
1162    + One
1163    + Zero
1164    + crate::MaybeSendSync
1165    + crate::simd::MaybeSimdOps
1166    + crate::simd::MaybeSimdProduct
1167    + crate::simd::MaybeSimdSumSquares
1168{
1169    fn reduce_sum(lhs: Self, rhs: Self) -> Self;
1170    fn reduce_product(lhs: Self, rhs: Self) -> Self;
1171    /// Identity of [`ReduceOp::Max`]; unreachable for dtypes the plan rejects.
1172    fn max_identity() -> Self;
1173    /// Identity of [`ReduceOp::Min`]; unreachable for dtypes the plan rejects.
1174    fn min_identity() -> Self;
1175    fn reduce_max(lhs: Self, rhs: Self) -> Self;
1176    fn reduce_min(lhs: Self, rhs: Self) -> Self;
1177}
1178
1179macro_rules! impl_float_erased_reduce_scalar {
1180    ($($ty:ty),* $(,)?) => {
1181        $(
1182            impl ErasedReduceScalar for $ty {
1183                #[inline(always)]
1184                fn reduce_sum(lhs: Self, rhs: Self) -> Self {
1185                    lhs + rhs
1186                }
1187
1188                #[inline(always)]
1189                fn reduce_product(lhs: Self, rhs: Self) -> Self {
1190                    lhs * rhs
1191                }
1192
1193                #[inline(always)]
1194                fn max_identity() -> Self {
1195                    <$ty>::NEG_INFINITY
1196                }
1197
1198                #[inline(always)]
1199                fn min_identity() -> Self {
1200                    <$ty>::INFINITY
1201                }
1202
1203                #[inline(always)]
1204                fn reduce_max(lhs: Self, rhs: Self) -> Self {
1205                    // Select form: both sides are evaluated and the NaN test
1206                    // picks the result, so lane loops lower to compare and
1207                    // blend instead of a data dependent branch.
1208                    // A NaN `lhs` is kept by the second test and a NaN
1209                    // `rhs` fails the comparison, so NaN propagates from
1210                    // either side. Equal values (including +0.0 and -0.0)
1211                    // return `rhs`.
1212                    if (lhs > rhs) | lhs.is_nan() {
1213                        lhs
1214                    } else {
1215                        rhs
1216                    }
1217                }
1218
1219                #[inline(always)]
1220                fn reduce_min(lhs: Self, rhs: Self) -> Self {
1221                    // Select form: both sides are evaluated and the NaN test
1222                    // picks the result, so lane loops lower to compare and
1223                    // blend instead of a data dependent branch.
1224                    // A NaN `lhs` is kept by the second test and a NaN
1225                    // `rhs` fails the comparison, so NaN propagates from
1226                    // either side. Equal values (including +0.0 and -0.0)
1227                    // return `rhs`.
1228                    if (lhs < rhs) | lhs.is_nan() {
1229                        lhs
1230                    } else {
1231                        rhs
1232                    }
1233                }
1234            }
1235        )*
1236    };
1237}
1238
1239macro_rules! impl_complex_erased_reduce_scalar {
1240    ($($ty:ty),* $(,)?) => {
1241        $(
1242            impl ErasedReduceScalar for $ty {
1243                #[inline(always)]
1244                fn reduce_sum(lhs: Self, rhs: Self) -> Self {
1245                    lhs + rhs
1246                }
1247
1248                #[inline(always)]
1249                fn reduce_product(lhs: Self, rhs: Self) -> Self {
1250                    lhs * rhs
1251                }
1252
1253                fn max_identity() -> Self {
1254                    // INVARIANT: check_reduce_op_dtype rejects complex Max at compile time.
1255                    unreachable!("complex max reduction is rejected at plan compile")
1256                }
1257
1258                fn min_identity() -> Self {
1259                    // INVARIANT: check_reduce_op_dtype rejects complex Min at compile time.
1260                    unreachable!("complex min reduction is rejected at plan compile")
1261                }
1262
1263                fn reduce_max(_lhs: Self, _rhs: Self) -> Self {
1264                    // INVARIANT: check_reduce_op_dtype rejects complex Max at compile time.
1265                    unreachable!("complex max reduction is rejected at plan compile")
1266                }
1267
1268                fn reduce_min(_lhs: Self, _rhs: Self) -> Self {
1269                    // INVARIANT: check_reduce_op_dtype rejects complex Min at compile time.
1270                    unreachable!("complex min reduction is rejected at plan compile")
1271                }
1272            }
1273        )*
1274    };
1275}
1276
1277macro_rules! impl_wrapping_erased_reduce_scalar {
1278    ($($ty:ty),* $(,)?) => {
1279        $(
1280            impl ErasedReduceScalar for $ty {
1281                #[inline(always)]
1282                fn reduce_sum(lhs: Self, rhs: Self) -> Self {
1283                    lhs.wrapping_add(rhs)
1284                }
1285
1286                #[inline(always)]
1287                fn reduce_product(lhs: Self, rhs: Self) -> Self {
1288                    lhs.wrapping_mul(rhs)
1289                }
1290
1291                #[inline(always)]
1292                fn max_identity() -> Self {
1293                    <$ty>::MIN
1294                }
1295
1296                #[inline(always)]
1297                fn min_identity() -> Self {
1298                    <$ty>::MAX
1299                }
1300
1301                #[inline(always)]
1302                fn reduce_max(lhs: Self, rhs: Self) -> Self {
1303                    lhs.max(rhs)
1304                }
1305
1306                #[inline(always)]
1307                fn reduce_min(lhs: Self, rhs: Self) -> Self {
1308                    lhs.min(rhs)
1309                }
1310            }
1311        )*
1312    };
1313}
1314
1315impl_float_erased_reduce_scalar!(f32, f64);
1316
1317impl_complex_erased_reduce_scalar!(Complex32, Complex64);
1318
1319impl_wrapping_erased_reduce_scalar!(i32, i64);
1320
1321fn validate_unique_axes(axes: &[usize], rank: usize) -> Result<()> {
1322    let mut seen = vec![false; rank];
1323    for &axis in axes {
1324        if axis >= rank {
1325            return Err(StridedError::InvalidAxis { axis, rank });
1326        }
1327        if seen[axis] {
1328            return Err(StridedError::InvalidAxis { axis, rank });
1329        }
1330        seen[axis] = true;
1331    }
1332    Ok(())
1333}
1334
1335struct ReduceOuterCursor<'a> {
1336    axes: &'a [ReduceOuterAxis],
1337    coords: CoordScratch,
1338    source_offset: isize,
1339    dest_offset: isize,
1340}
1341impl<'a> ReduceOuterCursor<'a> {
1342    fn decode(
1343        mut linear: usize,
1344        source_base: isize,
1345        dest_base: isize,
1346        axes: &'a [ReduceOuterAxis],
1347    ) -> Result<Self> {
1348        let mut coords = CoordScratch::new(axes.len());
1349        let mut source_offset = source_base;
1350        let mut dest_offset = dest_base;
1351        for (coord, axis) in coords.as_mut_slice().iter_mut().zip(axes) {
1352            // INVARIANT: a non-empty destination domain implies every outer
1353            // extent is nonzero before decode; compile-time span checks make
1354            // these one-time checked additions represent valid layout offsets.
1355            debug_assert!(axis.extent != 0);
1356            *coord = linear % axis.extent;
1357            linear /= axis.extent;
1358            source_offset = checked_offset_add(source_offset, axis.source_step, *coord)?;
1359            dest_offset = checked_offset_add(dest_offset, axis.dest_step, *coord)?;
1360        }
1361        Ok(Self {
1362            axes,
1363            coords,
1364            source_offset,
1365            dest_offset,
1366        })
1367    }
1368
1369    /// Coordinate along the first (fastest) compressed outer axis.
1370    #[inline]
1371    fn leading_coord(&self) -> usize {
1372        self.coords.first()
1373    }
1374
1375    #[inline]
1376    fn advance(&mut self) {
1377        // INVARIANT: compile_axes checked every signed step and reset delta;
1378        // descriptor validation plus exact layout equality proves each cursor
1379        // state is a reachable source/destination offset.
1380        for (coord, axis) in self.coords.as_mut_slice().iter_mut().zip(self.axes) {
1381            let next = *coord + 1;
1382            if next < axis.extent {
1383                *coord = next;
1384                self.source_offset += axis.source_step;
1385                self.dest_offset += axis.dest_step;
1386                return;
1387            }
1388            *coord = 0;
1389            self.source_offset += axis.source_reset;
1390            self.dest_offset += axis.dest_reset;
1391        }
1392    }
1393}
1394struct ReduceInnerCursor<'a> {
1395    axes: &'a [ReduceInnerAxis],
1396    coords: CoordScratch,
1397    source_offset: isize,
1398}
1399impl<'a> ReduceInnerCursor<'a> {
1400    fn new(source_base: isize, axes: &'a [ReduceInnerAxis]) -> Self {
1401        Self {
1402            axes,
1403            coords: CoordScratch::new(axes.len()),
1404            source_offset: source_base,
1405        }
1406    }
1407
1408    #[inline]
1409    fn reset(&mut self, source_base: isize) {
1410        self.coords.as_mut_slice().fill(0);
1411        self.source_offset = source_base;
1412    }
1413
1414    #[inline]
1415    fn advance(&mut self) {
1416        // INVARIANT: compile_axes checked every signed source step and reset
1417        // delta, and the validated descriptor/layout chain proves each value
1418        // offset is reachable from the current outer source base.
1419        for (coord, axis) in self.coords.as_mut_slice().iter_mut().zip(self.axes) {
1420            let next = *coord + 1;
1421            if next < axis.extent {
1422                *coord = next;
1423                self.source_offset += axis.source_step;
1424                return;
1425            }
1426            *coord = 0;
1427            self.source_offset += axis.source_reset;
1428        }
1429    }
1430}
1431fn check_reduce_layout_offset_arithmetic(dims: &[usize], strides: &[isize]) -> Result<()> {
1432    if dims.len() != strides.len() {
1433        return Err(StridedError::StrideLengthMismatch);
1434    }
1435    let mut min_offset = 0isize;
1436    let mut max_offset = 0isize;
1437    for (&dim, &stride) in dims.iter().zip(strides) {
1438        let last =
1439            isize::try_from(dim.saturating_sub(1)).map_err(|_| StridedError::OffsetOverflow)?;
1440        let extent = stride
1441            .checked_mul(last)
1442            .ok_or(StridedError::OffsetOverflow)?;
1443        if extent < 0 {
1444            min_offset = min_offset
1445                .checked_add(extent)
1446                .ok_or(StridedError::OffsetOverflow)?;
1447        } else {
1448            max_offset = max_offset
1449                .checked_add(extent)
1450                .ok_or(StridedError::OffsetOverflow)?;
1451        }
1452    }
1453    let _ = (min_offset, max_offset);
1454    Ok(())
1455}
1456fn compress_reduce_outer_axes(axes: Vec<ReduceOuterAxis>) -> Result<Vec<ReduceOuterAxis>> {
1457    let mut compressed: Vec<ReduceOuterAxis> = Vec::with_capacity(axes.len());
1458    for axis in axes {
1459        if let Some(previous) = compressed.last_mut() {
1460            let previous_extent =
1461                isize::try_from(previous.extent).map_err(|_| StridedError::OffsetOverflow)?;
1462            let expected_source = previous
1463                .source_step
1464                .checked_mul(previous_extent)
1465                .ok_or(StridedError::OffsetOverflow)?;
1466            let expected_dest = previous
1467                .dest_step
1468                .checked_mul(previous_extent)
1469                .ok_or(StridedError::OffsetOverflow)?;
1470            if axis.source_step == expected_source && axis.dest_step == expected_dest {
1471                let fused_extent = previous
1472                    .extent
1473                    .checked_mul(axis.extent)
1474                    .ok_or(StridedError::OffsetOverflow)?;
1475                previous.extent = fused_extent;
1476                previous.source_reset = checked_reduce_reset(fused_extent, previous.source_step)?;
1477                previous.dest_reset = checked_reduce_reset(fused_extent, previous.dest_step)?;
1478                continue;
1479            }
1480        }
1481        compressed.push(axis);
1482    }
1483    Ok(compressed)
1484}
1485fn compress_reduce_inner_axes(axes: Vec<ReduceInnerAxis>) -> Result<Vec<ReduceInnerAxis>> {
1486    let mut compressed: Vec<ReduceInnerAxis> = Vec::with_capacity(axes.len());
1487    for axis in axes {
1488        if let Some(previous) = compressed.last_mut() {
1489            let previous_extent =
1490                isize::try_from(previous.extent).map_err(|_| StridedError::OffsetOverflow)?;
1491            let expected_source = previous
1492                .source_step
1493                .checked_mul(previous_extent)
1494                .ok_or(StridedError::OffsetOverflow)?;
1495            if axis.source_step == expected_source {
1496                let fused_extent = previous
1497                    .extent
1498                    .checked_mul(axis.extent)
1499                    .ok_or(StridedError::OffsetOverflow)?;
1500                previous.extent = fused_extent;
1501                previous.source_reset = checked_reduce_reset(fused_extent, previous.source_step)?;
1502                continue;
1503            }
1504        }
1505        compressed.push(axis);
1506    }
1507    Ok(compressed)
1508}
1509fn checked_reduce_reset(extent: usize, stride: isize) -> Result<isize> {
1510    if extent == 0 {
1511        return Ok(0);
1512    }
1513    let last = isize::try_from(extent - 1).map_err(|_| StridedError::OffsetOverflow)?;
1514    stride
1515        .checked_mul(last)
1516        .and_then(isize::checked_neg)
1517        .ok_or(StridedError::OffsetOverflow)
1518}
1519fn checked_offset_add(base: isize, stride: isize, coord: usize) -> Result<isize> {
1520    let coord = isize::try_from(coord).map_err(|_| StridedError::OffsetOverflow)?;
1521    let scaled = stride
1522        .checked_mul(coord)
1523        .ok_or(StridedError::OffsetOverflow)?;
1524    base.checked_add(scaled).ok_or(StridedError::OffsetOverflow)
1525}
1526
1527struct CoordScratch {
1528    inline: [usize; RAW_FUSED_RANK_LIMIT],
1529    heap: Option<Vec<usize>>,
1530    len: usize,
1531}
1532
1533impl CoordScratch {
1534    fn new(len: usize) -> Self {
1535        if len <= RAW_FUSED_RANK_LIMIT {
1536            Self {
1537                inline: [0; RAW_FUSED_RANK_LIMIT],
1538                heap: None,
1539                len,
1540            }
1541        } else {
1542            Self {
1543                inline: [0; RAW_FUSED_RANK_LIMIT],
1544                heap: Some(vec![0; len]),
1545                len,
1546            }
1547        }
1548    }
1549
1550    #[inline]
1551    fn first(&self) -> usize {
1552        match &self.heap {
1553            Some(heap) => heap.first().copied().unwrap_or(0),
1554            None => self.inline[0],
1555        }
1556    }
1557
1558    fn as_mut_slice(&mut self) -> &mut [usize] {
1559        match &mut self.heap {
1560            Some(heap) => heap,
1561            None => &mut self.inline[..self.len],
1562        }
1563    }
1564}