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}
229
230/// Dtype-erased reduction wrapper.
231///
232/// This is the erased replay boundary for full-tensor scalar reductions and
233/// axis reductions with a fixed output layout. It supports only operations with
234/// an unambiguous identity value in the selected dtype.
235#[derive(Clone, Debug)]
236pub struct ErasedReducePlan {
237    dtype: KernelDType,
238    op: ReduceOp,
239    layout: ReduceLayout,
240}
241
242#[derive(Clone, Debug)]
243enum ReduceLayout {
244    Full {
245        dims: Vec<usize>,
246        src_strides: Vec<isize>,
247    },
248    Axes {
249        src_dims: Vec<usize>,
250        src_strides: Vec<isize>,
251        dest_dims: Vec<usize>,
252        dest_strides: Vec<isize>,
253        axes: Vec<usize>,
254        kept_axes: Vec<usize>,
255        outer_axes: Vec<ReduceOuterAxis>,
256        inner_axes: Vec<ReduceInnerAxis>,
257        dest_total: usize,
258        reduce_total: usize,
259    },
260}
261
262#[derive(Clone, Copy, Debug)]
263struct ReduceOuterAxis {
264    extent: usize,
265    source_step: isize,
266    source_reset: isize,
267    dest_step: isize,
268    dest_reset: isize,
269}
270#[derive(Clone, Copy, Debug)]
271struct ReduceInnerAxis {
272    extent: usize,
273    source_step: isize,
274    source_reset: isize,
275}
276impl ReduceLayout {
277    fn src_dims(&self) -> &[usize] {
278        match self {
279            Self::Full { dims, .. } => dims,
280            Self::Axes { src_dims, .. } => src_dims,
281        }
282    }
283
284    fn src_strides(&self) -> &[isize] {
285        match self {
286            Self::Full { src_strides, .. } | Self::Axes { src_strides, .. } => src_strides,
287        }
288    }
289
290    fn check_src_layout(&self, src: &ErasedRawStridedRef<'_>) -> Result<()> {
291        if src.dims() != self.src_dims() || src.strides() != self.src_strides() {
292            return Err(StridedError::PlanLayoutMismatch);
293        }
294        Ok(())
295    }
296}
297
298#[derive(Clone, Copy, Debug)]
299struct AxesLayout<'a> {
300    src_dims: &'a [usize],
301    axes: &'a [usize],
302    kept_axes: &'a [usize],
303    outer_axes: &'a [ReduceOuterAxis],
304    inner_axes: &'a [ReduceInnerAxis],
305    dest_total: usize,
306    reduce_total: usize,
307}
308
309impl ErasedReducePlan {
310    /// Validate and store a full-reduction plan for one dtype and source layout.
311    pub fn compile(
312        dtype: KernelDType,
313        op: ReduceOp,
314        dims: &[usize],
315        src_strides: &[isize],
316    ) -> Result<Self> {
317        check_reduce_op_dtype(dtype, op)?;
318        if dims.len() != src_strides.len() {
319            return Err(StridedError::StrideLengthMismatch);
320        }
321        checked_total_len(dims)?;
322        Ok(Self {
323            dtype,
324            op,
325            layout: ReduceLayout::Full {
326                dims: dims.to_vec(),
327                src_strides: src_strides.to_vec(),
328            },
329        })
330    }
331
332    /// Validate and store an axis-reduction plan for one dtype and fixed source/output layouts.
333    ///
334    /// `axes` names the source axes reduced away. Output dimensions must be the
335    /// remaining source dimensions in source-axis order. When all axes are
336    /// reduced, any output layout with exactly one reachable element is accepted.
337    #[allow(clippy::too_many_arguments)]
338    pub fn compile_axes(
339        dtype: KernelDType,
340        op: ReduceOp,
341        src_dims: &[usize],
342        src_strides: &[isize],
343        dest_dims: &[usize],
344        dest_strides: &[isize],
345        axes: &[usize],
346    ) -> Result<Self> {
347        check_reduce_op_dtype(dtype, op)?;
348        if src_dims.len() != src_strides.len() || dest_dims.len() != dest_strides.len() {
349            return Err(StridedError::StrideLengthMismatch);
350        }
351        checked_total_len(src_dims)?;
352        check_reduce_layout_offset_arithmetic(src_dims, src_strides)?;
353        let dest_total = checked_total_len(dest_dims)?;
354        check_reduce_layout_offset_arithmetic(dest_dims, dest_strides)?;
355        if !crate::layout_check::is_injective_layout(dest_dims, dest_strides) {
356            return Err(StridedError::NonInjectiveOutputLayout);
357        }
358        validate_unique_axes(axes, src_dims.len())?;
359
360        let kept_axes: Vec<usize> = (0..src_dims.len())
361            .filter(|axis| !axes.contains(axis))
362            .collect();
363        let expected_dest_dims: Vec<usize> = kept_axes.iter().map(|&axis| src_dims[axis]).collect();
364        if expected_dest_dims.is_empty() {
365            if dest_total != 1 {
366                return Err(StridedError::ShapeMismatch(
367                    dest_dims.to_vec(),
368                    expected_dest_dims,
369                ));
370            }
371        } else if dest_dims != expected_dest_dims.as_slice() {
372            return Err(StridedError::ShapeMismatch(
373                dest_dims.to_vec(),
374                expected_dest_dims,
375            ));
376        }
377
378        let reduce_total = axes
379            .iter()
380            .try_fold(1usize, |total, &axis| total.checked_mul(src_dims[axis]))
381            .ok_or(StridedError::OffsetOverflow)?;
382        let outer_axes = compress_reduce_outer_axes(
383            kept_axes
384                .iter()
385                .enumerate()
386                .map(|(dest_axis, &src_axis)| {
387                    let extent = src_dims[src_axis];
388                    Ok(ReduceOuterAxis {
389                        extent,
390                        source_step: src_strides[src_axis],
391                        source_reset: checked_reduce_reset(extent, src_strides[src_axis])?,
392                        dest_step: dest_strides[dest_axis],
393                        dest_reset: checked_reduce_reset(extent, dest_strides[dest_axis])?,
394                    })
395                })
396                .collect::<Result<Vec<_>>>()?,
397        )?;
398        let inner_axes = compress_reduce_inner_axes(
399            axes.iter()
400                .map(|&src_axis| {
401                    let extent = src_dims[src_axis];
402                    Ok(ReduceInnerAxis {
403                        extent,
404                        source_step: src_strides[src_axis],
405                        source_reset: checked_reduce_reset(extent, src_strides[src_axis])?,
406                    })
407                })
408                .collect::<Result<Vec<_>>>()?,
409        )?;
410        Ok(Self {
411            dtype,
412            op,
413            layout: ReduceLayout::Axes {
414                src_dims: src_dims.to_vec(),
415                src_strides: src_strides.to_vec(),
416                dest_dims: dest_dims.to_vec(),
417                dest_strides: dest_strides.to_vec(),
418                axes: axes.to_vec(),
419                kept_axes,
420                outer_axes,
421                inner_axes,
422                dest_total,
423                reduce_total,
424            },
425        })
426    }
427
428    #[inline]
429    pub fn dtype(&self) -> KernelDType {
430        self.dtype
431    }
432
433    #[inline]
434    pub fn op(&self) -> ReduceOp {
435        self.op
436    }
437
438    /// Execute the reduction into an erased output descriptor.
439    pub fn execute(
440        &self,
441        ctx: &ExecContext,
442        dest: &mut ErasedRawStridedMut<'_>,
443        src: &ErasedRawStridedRef<'_>,
444    ) -> Result<()> {
445        check_dtype(self.dtype, dest.dtype())?;
446        check_dtype(self.dtype, src.dtype())?;
447        self.layout.check_src_layout(src)?;
448        match &self.layout {
449            ReduceLayout::Full { .. } => {
450                let dest_len = checked_total_len(dest.dims())?;
451                if dest_len != 1 {
452                    return Err(StridedError::RankMismatch(dest_len, 1));
453                }
454            }
455            ReduceLayout::Axes {
456                dest_dims,
457                dest_strides,
458                ..
459            } => {
460                if dest.dims() != dest_dims.as_slice() || dest.strides() != dest_strides.as_slice()
461                {
462                    return Err(StridedError::PlanLayoutMismatch);
463                }
464            }
465        }
466
467        let result = match self.dtype {
468            KernelDType::F32 => {
469                let mut writer = reduce_writer::<f32>(dest)?;
470                dispatch_reduce::<f32, _>(self.op, &self.layout, ctx, &mut writer, src)
471            }
472            KernelDType::F64 => {
473                let mut writer = reduce_writer::<f64>(dest)?;
474                dispatch_reduce::<f64, _>(self.op, &self.layout, ctx, &mut writer, src)
475            }
476            KernelDType::I32 => {
477                let mut writer = reduce_writer::<i32>(dest)?;
478                dispatch_reduce::<i32, _>(self.op, &self.layout, ctx, &mut writer, src)
479            }
480            KernelDType::I64 => {
481                let mut writer = reduce_writer::<i64>(dest)?;
482                dispatch_reduce::<i64, _>(self.op, &self.layout, ctx, &mut writer, src)
483            }
484            KernelDType::C32 => {
485                let mut writer = reduce_writer::<Complex32>(dest)?;
486                dispatch_reduce::<Complex32, _>(self.op, &self.layout, ctx, &mut writer, src)
487            }
488            KernelDType::C64 => {
489                let mut writer = reduce_writer::<Complex64>(dest)?;
490                dispatch_reduce::<Complex64, _>(self.op, &self.layout, ctx, &mut writer, src)
491            }
492            _ => Err(StridedError::UnsupportedDType {
493                dtype: self.dtype.label(),
494            }),
495        };
496        result
497    }
498
499    /// On success, every reachable destination slot is fully overwritten;
500    /// unreachable holes are neither read nor initialized. Validation errors
501    /// are returned before any destination write. A panic during execution
502    /// may leave a partially initialized `MaybeUninit` destination, which is
503    /// still safely droppable; no readable value is promised for unwritten
504    /// reachable slots.
505    pub fn execute_uninit(
506        &self,
507        ctx: &ExecContext,
508        dest: &mut ErasedRawStridedUninitMut<'_>,
509        src: &ErasedRawStridedPtr<'_>,
510    ) -> Result<()> {
511        check_dtype(self.dtype, dest.dtype())?;
512        check_dtype(self.dtype, src.dtype())?;
513        validate_uninit_no_overlap(dest, src, 0)?;
514        // SAFETY: the owning erased entry rejected all input/output overlap before conversion.
515        let src = unsafe { src.try_as_ref_after_no_overlap() }?;
516        self.layout.check_src_layout(&src)?;
517        match &self.layout {
518            ReduceLayout::Full { .. } => {
519                let total = checked_total_len(dest.dims())?;
520                if total != 1 {
521                    return Err(StridedError::RankMismatch(total, 1));
522                }
523            }
524            ReduceLayout::Axes {
525                dest_dims,
526                dest_strides,
527                ..
528            } => {
529                if dest.dims() != dest_dims.as_slice() || dest.strides() != dest_strides.as_slice()
530                {
531                    return Err(StridedError::PlanLayoutMismatch);
532                }
533            }
534        }
535        macro_rules! run {
536            ($ty:ty) => {{
537                let mut writer = reduce_uninit_writer::<$ty>(dest)?;
538                dispatch_reduce::<$ty, _>(self.op, &self.layout, ctx, &mut writer, &src)
539            }};
540        }
541        match self.dtype {
542            KernelDType::F32 => run!(f32),
543            KernelDType::F64 => run!(f64),
544            KernelDType::I32 => run!(i32),
545            KernelDType::I64 => run!(i64),
546            KernelDType::C32 => run!(Complex32),
547            KernelDType::C64 => run!(Complex64),
548            _ => Err(StridedError::UnsupportedDType {
549                dtype: self.dtype.label(),
550            }),
551        }
552    }
553}
554
555fn reduce_writer<'a, T>(dest: &'a mut ErasedRawStridedMut<'_>) -> Result<RawReduceWriter<'a, T>>
556where
557    T: KernelStorageElement,
558{
559    let offset = dest.offset();
560    let data = dest.data_as_mut::<T>()?;
561    let ptr = data.as_mut_ptr();
562    let extent = data.len();
563    Ok(RawReduceWriter {
564        ptr,
565        extent,
566        offset,
567        _marker: core::marker::PhantomData,
568    })
569}
570
571fn reduce_uninit_writer<'a, T>(
572    dest: &'a mut ErasedRawStridedUninitMut<'_>,
573) -> Result<RawReduceWriter<'a, T>>
574where
575    T: KernelStorageElement,
576{
577    let offset = dest.offset();
578    let data = dest.data_as_uninit_mut::<T>()?;
579    let ptr = data.as_mut_ptr().cast::<T>();
580    let extent = data.len();
581    Ok(RawReduceWriter {
582        ptr,
583        extent,
584        offset,
585        _marker: core::marker::PhantomData,
586    })
587}
588
589fn check_reduce_dtype(dtype: KernelDType) -> Result<()> {
590    match dtype {
591        KernelDType::F32
592        | KernelDType::F64
593        | KernelDType::I32
594        | KernelDType::I64
595        | KernelDType::C32
596        | KernelDType::C64 => Ok(()),
597        _ => Err(StridedError::UnsupportedDType {
598            dtype: dtype.label(),
599        }),
600    }
601}
602
603fn check_reduce_op_dtype(dtype: KernelDType, op: ReduceOp) -> Result<()> {
604    if op == ReduceOp::SumSquares && !matches!(dtype, KernelDType::F32 | KernelDType::F64) {
605        return Err(StridedError::UnsupportedDType {
606            dtype: dtype.label(),
607        });
608    }
609    check_reduce_dtype(dtype)
610}
611
612fn checked_total_len(dims: &[usize]) -> Result<usize> {
613    if dims.is_empty() {
614        return Ok(1);
615    }
616    dims.iter()
617        .try_fold(1usize, |acc, &dim| acc.checked_mul(dim))
618        .ok_or(StridedError::OffsetOverflow)
619}
620
621fn execute_copy<T>(
622    plan: &CopyPlan,
623    dest: &mut ErasedRawStridedMut<'_>,
624    src: &ErasedRawStridedRef<'_>,
625) -> Result<()>
626where
627    T: Copy + crate::MaybeSendSync + KernelStorageElement,
628{
629    let source_data = src.data_as::<T>()?;
630    let dest_dims = dest.dims();
631    let dest_strides = dest.strides();
632    let dest_offset = dest.offset();
633    let dest_data = dest.data_as_mut::<T>()?;
634    let source = unsafe {
635        RawStridedRef::new_unchecked(source_data, src.dims(), src.strides(), src.offset())
636    };
637    let mut dest =
638        unsafe { RawStridedMut::new_unchecked(dest_data, dest_dims, dest_strides, dest_offset) };
639    plan.execute(&mut dest, &source)
640}
641
642fn execute_concatenate<T>(
643    plan: &ConcatenatePlan,
644    dest: &mut ErasedRawStridedMut<'_>,
645    inputs: &[ErasedRawStridedRef<'_>],
646) -> Result<()>
647where
648    T: Copy + crate::MaybeSendSync + KernelStorageElement,
649{
650    if inputs.len() != plan.input_count() {
651        return Err(StridedError::RankMismatch(inputs.len(), plan.input_count()));
652    }
653    let dest_dims = dest.dims();
654    let dest_strides = dest.strides();
655    let dest_offset = dest.offset();
656    let dest_data = dest.data_as_mut::<T>()?;
657    let mut dest_ref =
658        unsafe { RawStridedMut::new_unchecked(dest_data, dest_dims, dest_strides, dest_offset) };
659    plan.check_dest_layout(&dest_ref)?;
660
661    for (position, input) in inputs.iter().enumerate() {
662        let input_data = input.data_as::<T>()?;
663        let input_ref = unsafe {
664            RawStridedRef::new_unchecked(input_data, input.dims(), input.strides(), input.offset())
665        };
666        plan.check_input_layout(position, &input_ref)?;
667        plan.segment_offset(position, dest_offset)?;
668    }
669    for (position, input) in inputs.iter().enumerate() {
670        let input_data = input.data_as::<T>()?;
671        let input_ref = unsafe {
672            RawStridedRef::new_unchecked(input_data, input.dims(), input.strides(), input.offset())
673        };
674        plan.execute_segment(position, &mut dest_ref, &input_ref)?;
675    }
676    Ok(())
677}
678
679fn execute_concatenate_uninit<T>(
680    plan: &ConcatenatePlan,
681    dest: &mut ErasedRawStridedUninitMut<'_>,
682    inputs: &[ErasedRawStridedPtr<'_>],
683) -> Result<()>
684where
685    T: Copy + crate::MaybeSendSync + KernelStorageElement,
686{
687    let dest_dims = dest.dims();
688    let dest_strides = dest.strides();
689    let dest_offset = dest.offset();
690    let dest_data = dest.data_as_uninit_mut::<T>()?;
691    let mut dest_ref =
692        unsafe { RawStridedMut::new_unchecked(dest_data, dest_dims, dest_strides, dest_offset) };
693    plan.check_dest_layout(&dest_ref)?;
694
695    for (position, input) in inputs.iter().enumerate() {
696        // SAFETY: input/output overlap was rejected before forming references.
697        let input = unsafe { input.try_as_ref_after_no_overlap() }?;
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.check_input_layout(position, &input_ref)?;
703        plan.segment_offset(position, dest_offset)?;
704    }
705    for (position, input) in inputs.iter().enumerate() {
706        // SAFETY: input/output overlap was rejected before forming references.
707        let input = unsafe { input.try_as_ref_after_no_overlap() }?;
708        let input_data = input.data_as::<T>()?;
709        let input_ref = unsafe {
710            RawStridedRef::new_unchecked(input_data, input.dims(), input.strides(), input.offset())
711        };
712        plan.execute_segment_uninit(position, &mut dest_ref, &input_ref)?;
713    }
714    Ok(())
715}
716
717fn execute_reduce<T, W>(
718    op: ReduceOp,
719    ctx: &ExecContext,
720    dest: &mut W,
721    src: &ErasedRawStridedRef<'_>,
722) -> Result<()>
723where
724    T: ErasedReduceScalar,
725    W: ReduceWriter<T>,
726{
727    let use_serial = ctx.is_serial()
728        || ctx
729            .max_threads_limit()
730            .is_some_and(|max_threads| max_threads.get() == 1);
731    let value = if use_serial {
732        if let Some(value) = reduce_contiguous_serial(op, src) {
733            value
734        } else {
735            let source = erased_view::<T>(src)?;
736            crate::reduce_view::reduce_serial(
737                &source,
738                |value| reduce_map_value(op, value),
739                |a, b| reduce_values(op, a, b),
740                reduce_identity(op),
741            )?
742        }
743    } else {
744        let source = erased_view::<T>(src)?;
745        ctx.run(|| {
746            crate::reduce(
747                &source,
748                |value| reduce_map_value(op, value),
749                |a, b| reduce_values(op, a, b),
750                reduce_identity(op),
751            )
752        })?
753    };
754
755    // SAFETY: validated rank-zero destination layout proves the offset.
756    unsafe { dest.write_at(dest.offset(), value) };
757    Ok(())
758}
759
760fn reduce_contiguous_serial<T>(op: ReduceOp, src: &ErasedRawStridedRef<'_>) -> Option<T>
761where
762    T: ErasedReduceScalar,
763{
764    crate::kernel::same_contiguous_layout(src.dims(), &[src.strides()])?;
765    let len = checked_total_len(src.dims()).ok()?;
766    if len == 0 {
767        return Some(reduce_identity(op));
768    }
769
770    let source_data = src.data_as::<T>().ok()?;
771    let start = usize::try_from(src.offset()).ok()?;
772    let end = start.checked_add(len)?;
773    let values = source_data.get(start..end)?;
774    Some(match op {
775        ReduceOp::Sum => T::try_simd_sum(values)
776            .unwrap_or_else(|| reduce_contiguous_lanes(values, T::zero(), T::reduce_sum)),
777        ReduceOp::Product => T::try_simd_product(values)
778            .unwrap_or_else(|| reduce_contiguous_lanes(values, T::one(), T::reduce_product)),
779        ReduceOp::SumSquares => T::try_simd_sum_squares(values).unwrap_or_else(|| {
780            reduce_contiguous_mapped_lanes(
781                values,
782                T::zero(),
783                |value| T::reduce_product(value, value),
784                T::reduce_sum,
785            )
786        }),
787    })
788}
789
790#[inline]
791fn reduce_contiguous_lanes<T>(values: &[T], identity: T, combine: impl Fn(T, T) -> T) -> T
792where
793    T: Copy,
794{
795    reduce_contiguous_mapped_lanes(values, identity, |value| value, combine)
796}
797
798#[inline]
799fn reduce_contiguous_mapped_lanes<T>(
800    values: &[T],
801    identity: T,
802    map: impl Fn(T) -> T,
803    combine: impl Fn(T, T) -> T,
804) -> T
805where
806    T: Copy,
807{
808    let mut lanes = [identity; SERIAL_REDUCE_LANES];
809    let mut chunks = values.chunks_exact(SERIAL_REDUCE_LANES);
810    for chunk in chunks.by_ref() {
811        for lane in 0..SERIAL_REDUCE_LANES {
812            lanes[lane] = combine(lanes[lane], map(chunk[lane]));
813        }
814    }
815    for (lane, &value) in chunks.remainder().iter().enumerate() {
816        lanes[lane] = combine(lanes[lane], map(value));
817    }
818    lanes.into_iter().fold(identity, combine)
819}
820
821fn dispatch_reduce<T, W>(
822    op: ReduceOp,
823    layout: &ReduceLayout,
824    ctx: &ExecContext,
825    dest: &mut W,
826    src: &ErasedRawStridedRef<'_>,
827) -> Result<()>
828where
829    T: ErasedReduceScalar,
830    W: ReduceWriter<T>,
831{
832    match layout {
833        ReduceLayout::Full { .. } => execute_reduce::<T, W>(op, ctx, dest, src),
834        ReduceLayout::Axes {
835            src_dims,
836            axes,
837            kept_axes,
838            outer_axes,
839            inner_axes,
840            dest_total,
841            reduce_total,
842            ..
843        } => execute_reduce_axes::<T, W>(
844            op,
845            ctx,
846            dest,
847            src,
848            AxesLayout {
849                src_dims,
850                axes,
851                kept_axes,
852                outer_axes,
853                inner_axes,
854                dest_total: *dest_total,
855                reduce_total: *reduce_total,
856            },
857        ),
858    }
859}
860
861fn execute_reduce_axes<T, W>(
862    op: ReduceOp,
863    ctx: &ExecContext,
864    dest: &mut W,
865    src: &ErasedRawStridedRef<'_>,
866    layout: AxesLayout<'_>,
867) -> Result<()>
868where
869    T: ErasedReduceScalar,
870    W: ReduceWriter<T>,
871{
872    if layout.kept_axes.is_empty()
873        && layout.axes.len() == layout.src_dims.len()
874        && layout.dest_total == 1
875    {
876        return execute_reduce::<T, W>(op, ctx, dest, src);
877    }
878
879    if layout.dest_total == 0 {
880        return Ok(());
881    }
882
883    if layout.reduce_total == 0 {
884        if ctx.is_serial() {
885            execute_reduce_axes_identity_serial(op, dest, layout)
886        } else {
887            ctx.run(|| execute_reduce_axes_identity_policy(op, dest, layout))
888        }
889    } else if ctx.is_serial() {
890        execute_reduce_axes_serial::<T, W>(op, dest, src, layout)
891    } else {
892        ctx.run(|| execute_reduce_axes_policy::<T, W>(op, dest, src, layout))
893    }
894}
895
896fn execute_reduce_axes_policy<T, W>(
897    op: ReduceOp,
898    dest: &mut W,
899    src: &ErasedRawStridedRef<'_>,
900    layout: AxesLayout<'_>,
901) -> Result<()>
902where
903    T: ErasedReduceScalar,
904    W: ReduceWriter<T>,
905{
906    let source_data = src.data_as::<T>()?;
907    let dest_offset_base = dest.offset();
908    #[cfg(feature = "parallel")]
909    {
910        let nthreads = crate::threading::parallel_threads_for_len(layout.dest_total);
911        if nthreads > 1 {
912            return execute_reduce_axes_parallel(
913                op,
914                dest_offset_base,
915                dest,
916                src.offset(),
917                source_data,
918                layout,
919                nthreads,
920            );
921        }
922    }
923
924    execute_reduce_axes_serial_data(
925        op,
926        dest_offset_base,
927        dest,
928        src.offset(),
929        source_data,
930        layout,
931    )
932}
933
934fn execute_reduce_axes_serial<T, W>(
935    op: ReduceOp,
936    dest: &mut W,
937    src: &ErasedRawStridedRef<'_>,
938    layout: AxesLayout<'_>,
939) -> Result<()>
940where
941    T: ErasedReduceScalar,
942    W: ReduceWriter<T>,
943{
944    let source_data = src.data_as::<T>()?;
945    execute_reduce_axes_serial_data(op, dest.offset(), dest, src.offset(), source_data, layout)
946}
947
948fn execute_reduce_axes_serial_data<T, W>(
949    op: ReduceOp,
950    dest_offset_base: isize,
951    dest: &mut W,
952    source_offset_base: isize,
953    source_data: &[T],
954    layout: AxesLayout<'_>,
955) -> Result<()>
956where
957    T: ErasedReduceScalar,
958    W: ReduceWriter<T>,
959{
960    let mut outer =
961        ReduceOuterCursor::decode(0, source_offset_base, dest_offset_base, layout.outer_axes)?;
962    // INVARIANT: (1) compile_axes checked signed source/destination spans and
963    // every cursor step/reset, including -(extent-1)*stride; (2) raw input and
964    // output descriptors validated every reachable offset; (3) execute checked
965    // exact plan-layout equality before dispatch.
966    let reduce_inner = |inner: &mut ReduceInnerCursor<'_>| {
967        let mut acc = reduce_identity(op);
968        for value_index in 0..layout.reduce_total {
969            // SAFETY: the three-link layout invariant above proves each source
970            // cursor offset is within `source_data`.
971            let value = unsafe { *source_data.as_ptr().offset(inner.source_offset) };
972            acc = reduce_values(op, acc, reduce_map_value(op, value));
973            if value_index + 1 < layout.reduce_total {
974                inner.advance();
975            }
976        }
977        acc
978    };
979
980    if layout.inner_axes.len() <= RAW_FUSED_RANK_LIMIT {
981        for output in 0..layout.dest_total {
982            let mut inner = ReduceInnerCursor::new(outer.source_offset, layout.inner_axes);
983            let acc = reduce_inner(&mut inner);
984            // SAFETY: the three-link layout invariant above proves the destination
985            // cursor offset is an in-bounds logical output offset.
986            unsafe { dest.write_at(outer.dest_offset, acc) };
987            if output + 1 < layout.dest_total {
988                outer.advance();
989            }
990        }
991    } else {
992        let mut inner = ReduceInnerCursor::new(source_offset_base, layout.inner_axes);
993        for output in 0..layout.dest_total {
994            inner.reset(outer.source_offset);
995            let acc = reduce_inner(&mut inner);
996            // SAFETY: the three-link layout invariant above proves the destination
997            // cursor offset is an in-bounds logical output offset.
998            unsafe { dest.write_at(outer.dest_offset, acc) };
999            if output + 1 < layout.dest_total {
1000                outer.advance();
1001            }
1002        }
1003    }
1004    Ok(())
1005}
1006
1007fn execute_reduce_axes_identity_serial<T, W>(
1008    op: ReduceOp,
1009    dest: &mut W,
1010    layout: AxesLayout<'_>,
1011) -> Result<()>
1012where
1013    T: ErasedReduceScalar,
1014    W: ReduceWriter<T>,
1015{
1016    let mut outer = ReduceOuterCursor::decode(0, 0, dest.offset(), layout.outer_axes)?;
1017    for output in 0..layout.dest_total {
1018        // INVARIANT: (1) compile_axes checked the destination span and every
1019        // destination step/reset, including -(extent-1)*stride; (2) the raw
1020        // destination descriptor validated every reachable offset; (3) execute
1021        // checked exact plan-layout equality before dispatch.
1022        // SAFETY: the three-link layout invariant proves this destination
1023        // cursor offset is in bounds; the source pointer is intentionally never
1024        // formed for an empty reduction domain.
1025        unsafe { dest.write_at(outer.dest_offset, reduce_identity(op)) };
1026        if output + 1 < layout.dest_total {
1027            outer.advance();
1028        }
1029    }
1030    Ok(())
1031}
1032fn execute_reduce_axes_identity_policy<T, W>(
1033    op: ReduceOp,
1034    dest: &mut W,
1035    layout: AxesLayout<'_>,
1036) -> Result<()>
1037where
1038    T: ErasedReduceScalar,
1039    W: ReduceWriter<T>,
1040{
1041    #[cfg(feature = "parallel")]
1042    {
1043        let nthreads = crate::threading::parallel_threads_for_len(layout.dest_total);
1044        if nthreads > 1 {
1045            return execute_reduce_axes_identity_parallel(op, dest, layout, nthreads);
1046        }
1047    }
1048    execute_reduce_axes_identity_serial(op, dest, layout)
1049}
1050#[cfg(feature = "parallel")]
1051fn execute_reduce_axes_identity_parallel<T, W>(
1052    op: ReduceOp,
1053    dest: &mut W,
1054    layout: AxesLayout<'_>,
1055    nthreads: usize,
1056) -> Result<()>
1057where
1058    T: ErasedReduceScalar,
1059    W: ReduceWriter<T>,
1060{
1061    // SAFETY: the validated reduction writer owns the destination allocation.
1062    let dest_ptr = crate::threading::SendPtr(unsafe { dest.ptr() });
1063    let dest_offset_base = dest.offset();
1064    crate::threading::parallel_map_reduce(
1065        0..layout.dest_total,
1066        nthreads,
1067        &|range| {
1068            let range_end = range.end;
1069            let mut outer =
1070                ReduceOuterCursor::decode(range.start, 0, dest_offset_base, layout.outer_axes)?;
1071            let dest_ptr = dest_ptr.as_ptr();
1072            for output in range {
1073                // INVARIANT: (1) compile_axes checked the destination span and
1074                // every destination step/reset, including
1075                // -(extent-1)*stride; (2) the raw destination descriptor
1076                // validated every reachable pointer offset; (3) execute checked
1077                // exact plan-layout equality before dispatch.
1078                // SAFETY: the three-link layout invariant proves this
1079                // destination offset is in bounds; no source pointer is formed.
1080                unsafe {
1081                    dest_ptr
1082                        .offset(outer.dest_offset)
1083                        .write(reduce_identity(op))
1084                };
1085                if output + 1 < range_end {
1086                    outer.advance();
1087                }
1088            }
1089            Ok(())
1090        },
1091        &|left, right| left.and(right),
1092    )
1093}
1094#[cfg(feature = "parallel")]
1095fn execute_reduce_axes_parallel<T, W>(
1096    op: ReduceOp,
1097    dest_offset_base: isize,
1098    dest: &mut W,
1099    source_offset_base: isize,
1100    source_data: &[T],
1101    layout: AxesLayout<'_>,
1102    nthreads: usize,
1103) -> Result<()>
1104where
1105    T: ErasedReduceScalar,
1106    W: ReduceWriter<T>,
1107{
1108    // SAFETY: the validated reduction writer owns the destination allocation.
1109    let dest_ptr = crate::threading::SendPtr(unsafe { dest.ptr() });
1110    let source_ptr = crate::threading::SendPtr(source_data.as_ptr() as *mut T);
1111    crate::threading::parallel_map_reduce(
1112        0..layout.dest_total,
1113        nthreads,
1114        &|range| {
1115            let range_end = range.end;
1116            let mut outer = ReduceOuterCursor::decode(
1117                range.start,
1118                source_offset_base,
1119                dest_offset_base,
1120                layout.outer_axes,
1121            )?;
1122            let dest_ptr = dest_ptr.as_ptr();
1123            let source_ptr = source_ptr.as_const();
1124            // INVARIANT: (1) compile_axes checked signed source/destination
1125            // spans and every cursor step/reset, including -(extent-1)*stride;
1126            // (2) raw descriptors validated every reachable pointer offset;
1127            // (3) execute checked exact plan-layout equality before dispatch.
1128            let reduce_inner = |inner: &mut ReduceInnerCursor<'_>| {
1129                let mut acc = reduce_identity(op);
1130                for value_index in 0..layout.reduce_total {
1131                    // SAFETY: the three-link layout invariant above proves each
1132                    // source cursor offset is within the source allocation.
1133                    let value = unsafe { *source_ptr.offset(inner.source_offset) };
1134                    acc = reduce_values(op, acc, reduce_map_value(op, value));
1135                    if value_index + 1 < layout.reduce_total {
1136                        inner.advance();
1137                    }
1138                }
1139                acc
1140            };
1141
1142            if layout.inner_axes.len() <= RAW_FUSED_RANK_LIMIT {
1143                for output in range {
1144                    let mut inner = ReduceInnerCursor::new(outer.source_offset, layout.inner_axes);
1145                    let acc = reduce_inner(&mut inner);
1146                    // SAFETY: the three-link layout invariant above proves this
1147                    // destination cursor offset is in bounds.
1148                    unsafe { dest_ptr.offset(outer.dest_offset).write(acc) };
1149                    if output + 1 < range_end {
1150                        outer.advance();
1151                    }
1152                }
1153            } else {
1154                let mut inner = ReduceInnerCursor::new(outer.source_offset, layout.inner_axes);
1155                for output in range {
1156                    inner.reset(outer.source_offset);
1157                    let acc = reduce_inner(&mut inner);
1158                    // SAFETY: the three-link layout invariant above proves this
1159                    // destination cursor offset is in bounds.
1160                    unsafe { dest_ptr.offset(outer.dest_offset).write(acc) };
1161                    if output + 1 < range_end {
1162                        outer.advance();
1163                    }
1164                }
1165            }
1166            Ok(())
1167        },
1168        &|left, right| left.and(right),
1169    )
1170}
1171
1172#[inline]
1173fn reduce_identity<T>(op: ReduceOp) -> T
1174where
1175    T: One + Zero,
1176{
1177    match op {
1178        ReduceOp::Sum => T::zero(),
1179        ReduceOp::Product => T::one(),
1180        ReduceOp::SumSquares => T::zero(),
1181    }
1182}
1183
1184#[inline]
1185fn reduce_values<T>(op: ReduceOp, a: T, b: T) -> T
1186where
1187    T: ErasedReduceScalar,
1188{
1189    match op {
1190        ReduceOp::Sum => T::reduce_sum(a, b),
1191        ReduceOp::Product => T::reduce_product(a, b),
1192        ReduceOp::SumSquares => T::reduce_sum(a, b),
1193    }
1194}
1195
1196#[inline]
1197fn reduce_map_value<T>(op: ReduceOp, value: T) -> T
1198where
1199    T: ErasedReduceScalar,
1200{
1201    match op {
1202        ReduceOp::Sum | ReduceOp::Product => value,
1203        ReduceOp::SumSquares => T::reduce_product(value, value),
1204    }
1205}
1206
1207trait ErasedReduceScalar:
1208    KernelStorageElement
1209    + Copy
1210    + One
1211    + Zero
1212    + crate::MaybeSendSync
1213    + crate::simd::MaybeSimdOps
1214    + crate::simd::MaybeSimdProduct
1215    + crate::simd::MaybeSimdSumSquares
1216{
1217    fn reduce_sum(lhs: Self, rhs: Self) -> Self;
1218    fn reduce_product(lhs: Self, rhs: Self) -> Self;
1219}
1220
1221macro_rules! impl_default_erased_reduce_scalar {
1222    ($($ty:ty),* $(,)?) => {
1223        $(
1224            impl ErasedReduceScalar for $ty {
1225                #[inline(always)]
1226                fn reduce_sum(lhs: Self, rhs: Self) -> Self {
1227                    lhs + rhs
1228                }
1229
1230                #[inline(always)]
1231                fn reduce_product(lhs: Self, rhs: Self) -> Self {
1232                    lhs * rhs
1233                }
1234            }
1235        )*
1236    };
1237}
1238
1239macro_rules! impl_wrapping_erased_reduce_scalar {
1240    ($($ty:ty),* $(,)?) => {
1241        $(
1242            impl ErasedReduceScalar for $ty {
1243                #[inline(always)]
1244                fn reduce_sum(lhs: Self, rhs: Self) -> Self {
1245                    lhs.wrapping_add(rhs)
1246                }
1247
1248                #[inline(always)]
1249                fn reduce_product(lhs: Self, rhs: Self) -> Self {
1250                    lhs.wrapping_mul(rhs)
1251                }
1252            }
1253        )*
1254    };
1255}
1256
1257impl_default_erased_reduce_scalar!(f32, f64, Complex32, Complex64);
1258
1259impl_wrapping_erased_reduce_scalar!(i32, i64);
1260
1261fn validate_unique_axes(axes: &[usize], rank: usize) -> Result<()> {
1262    let mut seen = vec![false; rank];
1263    for &axis in axes {
1264        if axis >= rank {
1265            return Err(StridedError::InvalidAxis { axis, rank });
1266        }
1267        if seen[axis] {
1268            return Err(StridedError::InvalidAxis { axis, rank });
1269        }
1270        seen[axis] = true;
1271    }
1272    Ok(())
1273}
1274
1275struct ReduceOuterCursor<'a> {
1276    axes: &'a [ReduceOuterAxis],
1277    coords: CoordScratch,
1278    source_offset: isize,
1279    dest_offset: isize,
1280}
1281impl<'a> ReduceOuterCursor<'a> {
1282    fn decode(
1283        mut linear: usize,
1284        source_base: isize,
1285        dest_base: isize,
1286        axes: &'a [ReduceOuterAxis],
1287    ) -> Result<Self> {
1288        let mut coords = CoordScratch::new(axes.len());
1289        let mut source_offset = source_base;
1290        let mut dest_offset = dest_base;
1291        for (coord, axis) in coords.as_mut_slice().iter_mut().zip(axes) {
1292            // INVARIANT: a non-empty destination domain implies every outer
1293            // extent is nonzero before decode; compile-time span checks make
1294            // these one-time checked additions represent valid layout offsets.
1295            debug_assert!(axis.extent != 0);
1296            *coord = linear % axis.extent;
1297            linear /= axis.extent;
1298            source_offset = checked_offset_add(source_offset, axis.source_step, *coord)?;
1299            dest_offset = checked_offset_add(dest_offset, axis.dest_step, *coord)?;
1300        }
1301        Ok(Self {
1302            axes,
1303            coords,
1304            source_offset,
1305            dest_offset,
1306        })
1307    }
1308
1309    #[inline]
1310    fn advance(&mut self) {
1311        // INVARIANT: compile_axes checked every signed step and reset delta;
1312        // descriptor validation plus exact layout equality proves each cursor
1313        // state is a reachable source/destination offset.
1314        for (coord, axis) in self.coords.as_mut_slice().iter_mut().zip(self.axes) {
1315            let next = *coord + 1;
1316            if next < axis.extent {
1317                *coord = next;
1318                self.source_offset += axis.source_step;
1319                self.dest_offset += axis.dest_step;
1320                return;
1321            }
1322            *coord = 0;
1323            self.source_offset += axis.source_reset;
1324            self.dest_offset += axis.dest_reset;
1325        }
1326    }
1327}
1328struct ReduceInnerCursor<'a> {
1329    axes: &'a [ReduceInnerAxis],
1330    coords: CoordScratch,
1331    source_offset: isize,
1332}
1333impl<'a> ReduceInnerCursor<'a> {
1334    fn new(source_base: isize, axes: &'a [ReduceInnerAxis]) -> Self {
1335        Self {
1336            axes,
1337            coords: CoordScratch::new(axes.len()),
1338            source_offset: source_base,
1339        }
1340    }
1341
1342    #[inline]
1343    fn reset(&mut self, source_base: isize) {
1344        self.coords.as_mut_slice().fill(0);
1345        self.source_offset = source_base;
1346    }
1347
1348    #[inline]
1349    fn advance(&mut self) {
1350        // INVARIANT: compile_axes checked every signed source step and reset
1351        // delta, and the validated descriptor/layout chain proves each value
1352        // offset is reachable from the current outer source base.
1353        for (coord, axis) in self.coords.as_mut_slice().iter_mut().zip(self.axes) {
1354            let next = *coord + 1;
1355            if next < axis.extent {
1356                *coord = next;
1357                self.source_offset += axis.source_step;
1358                return;
1359            }
1360            *coord = 0;
1361            self.source_offset += axis.source_reset;
1362        }
1363    }
1364}
1365fn check_reduce_layout_offset_arithmetic(dims: &[usize], strides: &[isize]) -> Result<()> {
1366    if dims.len() != strides.len() {
1367        return Err(StridedError::StrideLengthMismatch);
1368    }
1369    let mut min_offset = 0isize;
1370    let mut max_offset = 0isize;
1371    for (&dim, &stride) in dims.iter().zip(strides) {
1372        let last =
1373            isize::try_from(dim.saturating_sub(1)).map_err(|_| StridedError::OffsetOverflow)?;
1374        let extent = stride
1375            .checked_mul(last)
1376            .ok_or(StridedError::OffsetOverflow)?;
1377        if extent < 0 {
1378            min_offset = min_offset
1379                .checked_add(extent)
1380                .ok_or(StridedError::OffsetOverflow)?;
1381        } else {
1382            max_offset = max_offset
1383                .checked_add(extent)
1384                .ok_or(StridedError::OffsetOverflow)?;
1385        }
1386    }
1387    let _ = (min_offset, max_offset);
1388    Ok(())
1389}
1390fn compress_reduce_outer_axes(axes: Vec<ReduceOuterAxis>) -> Result<Vec<ReduceOuterAxis>> {
1391    let mut compressed: Vec<ReduceOuterAxis> = Vec::with_capacity(axes.len());
1392    for axis in axes {
1393        if let Some(previous) = compressed.last_mut() {
1394            let previous_extent =
1395                isize::try_from(previous.extent).map_err(|_| StridedError::OffsetOverflow)?;
1396            let expected_source = previous
1397                .source_step
1398                .checked_mul(previous_extent)
1399                .ok_or(StridedError::OffsetOverflow)?;
1400            let expected_dest = previous
1401                .dest_step
1402                .checked_mul(previous_extent)
1403                .ok_or(StridedError::OffsetOverflow)?;
1404            if axis.source_step == expected_source && axis.dest_step == expected_dest {
1405                let fused_extent = previous
1406                    .extent
1407                    .checked_mul(axis.extent)
1408                    .ok_or(StridedError::OffsetOverflow)?;
1409                previous.extent = fused_extent;
1410                previous.source_reset = checked_reduce_reset(fused_extent, previous.source_step)?;
1411                previous.dest_reset = checked_reduce_reset(fused_extent, previous.dest_step)?;
1412                continue;
1413            }
1414        }
1415        compressed.push(axis);
1416    }
1417    Ok(compressed)
1418}
1419fn compress_reduce_inner_axes(axes: Vec<ReduceInnerAxis>) -> Result<Vec<ReduceInnerAxis>> {
1420    let mut compressed: Vec<ReduceInnerAxis> = Vec::with_capacity(axes.len());
1421    for axis in axes {
1422        if let Some(previous) = compressed.last_mut() {
1423            let previous_extent =
1424                isize::try_from(previous.extent).map_err(|_| StridedError::OffsetOverflow)?;
1425            let expected_source = previous
1426                .source_step
1427                .checked_mul(previous_extent)
1428                .ok_or(StridedError::OffsetOverflow)?;
1429            if axis.source_step == expected_source {
1430                let fused_extent = previous
1431                    .extent
1432                    .checked_mul(axis.extent)
1433                    .ok_or(StridedError::OffsetOverflow)?;
1434                previous.extent = fused_extent;
1435                previous.source_reset = checked_reduce_reset(fused_extent, previous.source_step)?;
1436                continue;
1437            }
1438        }
1439        compressed.push(axis);
1440    }
1441    Ok(compressed)
1442}
1443fn checked_reduce_reset(extent: usize, stride: isize) -> Result<isize> {
1444    if extent == 0 {
1445        return Ok(0);
1446    }
1447    let last = isize::try_from(extent - 1).map_err(|_| StridedError::OffsetOverflow)?;
1448    stride
1449        .checked_mul(last)
1450        .and_then(isize::checked_neg)
1451        .ok_or(StridedError::OffsetOverflow)
1452}
1453fn checked_offset_add(base: isize, stride: isize, coord: usize) -> Result<isize> {
1454    let coord = isize::try_from(coord).map_err(|_| StridedError::OffsetOverflow)?;
1455    let scaled = stride
1456        .checked_mul(coord)
1457        .ok_or(StridedError::OffsetOverflow)?;
1458    base.checked_add(scaled).ok_or(StridedError::OffsetOverflow)
1459}
1460
1461struct CoordScratch {
1462    inline: [usize; RAW_FUSED_RANK_LIMIT],
1463    heap: Option<Vec<usize>>,
1464    len: usize,
1465}
1466
1467impl CoordScratch {
1468    fn new(len: usize) -> Self {
1469        if len <= RAW_FUSED_RANK_LIMIT {
1470            Self {
1471                inline: [0; RAW_FUSED_RANK_LIMIT],
1472                heap: None,
1473                len,
1474            }
1475        } else {
1476            Self {
1477                inline: [0; RAW_FUSED_RANK_LIMIT],
1478                heap: Some(vec![0; len]),
1479                len,
1480            }
1481        }
1482    }
1483
1484    fn as_mut_slice(&mut self) -> &mut [usize] {
1485        match &mut self.heap {
1486            Some(heap) => heap,
1487            None => &mut self.inline[..self.len],
1488        }
1489    }
1490}