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