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