Skip to main content

strided_basic/
static_indexing_plan.rs

1//! Prepared static-indexing plans over raw strided value layouts.
2//!
3//! This module owns reusable static indexing traversal for downstream tensor
4//! runtimes. It keeps allocation, dtype policy, and tensor-level validation out
5//! of `strided-kernel`; callers provide already-owned output buffers and fixed
6//! raw descriptors.
7
8use core::mem::MaybeUninit;
9
10use crate::{CopyPlan, MaybeSendSync, RawStridedMut, RawStridedRef, Result, StridedError};
11
12#[cfg(feature = "parallel")]
13type AxisVec<T> = smallvec::SmallVec<[T; crate::RAW_FUSED_RANK_LIMIT]>;
14#[cfg(not(feature = "parallel"))]
15type AxisVec<T> = Vec<T>;
16
17/// A compiled static-slice traversal.
18///
19/// `compile` validates the fixed `starts`/`limits`/`slice_strides` contract and
20/// lowers replay to a strided copy from the corresponding source view.
21#[derive(Clone, Debug)]
22pub struct SlicePlan {
23    operand_dims: AxisVec<usize>,
24    operand_strides: AxisVec<isize>,
25    dest_dims: AxisVec<usize>,
26    dest_strides: AxisVec<isize>,
27    source_strides: AxisVec<isize>,
28    source_offset_delta: isize,
29    copy_plan: CopyPlan,
30}
31
32/// A compiled reverse traversal over selected axes.
33///
34/// Replay is a strided copy from a negative-stride source view.
35#[derive(Clone, Debug)]
36pub struct ReversePlan {
37    operand_dims: AxisVec<usize>,
38    operand_strides: AxisVec<isize>,
39    dest_strides: AxisVec<isize>,
40    source_strides: AxisVec<isize>,
41    source_offset_delta: isize,
42    copy_plan: CopyPlan,
43}
44
45/// A compiled pad traversal.
46///
47/// Replay first fills every destination element with the caller-provided fill
48/// scalar, then copies reachable input positions into the padded output.
49#[derive(Clone, Debug)]
50pub struct PadPlan {
51    operand_dims: AxisVec<usize>,
52    operand_strides: AxisVec<isize>,
53    dest_dims: AxisVec<usize>,
54    dest_strides: AxisVec<isize>,
55    edge_padding_low: AxisVec<i64>,
56    interior_step: AxisVec<i64>,
57    operand_total: usize,
58    dest_total: usize,
59    contiguous_dest_fill: bool,
60    contiguous_axis0_run: Option<ContiguousPadAxis0Run>,
61    generic_fill: PadFillCursor,
62    generic_copy: PadCopyCursor,
63}
64
65#[derive(Clone, Debug)]
66struct PadFillCursor {
67    steps: AxisVec<isize>,
68    resets: AxisVec<isize>,
69}
70
71#[derive(Clone, Debug)]
72struct PadCopyCursor {
73    shape: AxisVec<usize>,
74    source_base_delta: isize,
75    dest_base_delta: isize,
76    source_steps: AxisVec<isize>,
77    source_resets: AxisVec<isize>,
78    dest_steps: AxisVec<isize>,
79    dest_resets: AxisVec<isize>,
80    total: usize,
81}
82
83#[derive(Clone, Copy, Debug, Eq, PartialEq)]
84struct ContiguousPadAxis0Run {
85    operand_start: usize,
86    dest_start: usize,
87    len: usize,
88}
89
90/// A compiled multi-input concatenate traversal.
91///
92/// Each input segment is lowered to a prepared strided copy into the
93/// corresponding destination window.
94#[derive(Clone, Debug)]
95pub struct ConcatenatePlan {
96    input_dims: Vec<AxisVec<usize>>,
97    input_strides: Vec<AxisVec<isize>>,
98    dest_dims: AxisVec<usize>,
99    dest_strides: AxisVec<isize>,
100    dest_offset_deltas: Vec<isize>,
101    copy_plans: Vec<CopyPlan>,
102}
103
104impl SlicePlan {
105    /// Compile a static slice plan for one operand layout and destination layout.
106    #[allow(clippy::too_many_arguments)]
107    pub fn compile(
108        operand_dims: &[usize],
109        operand_strides: &[isize],
110        dest_dims: &[usize],
111        dest_strides: &[isize],
112        starts: &[usize],
113        limits: &[usize],
114        slice_strides: &[usize],
115    ) -> Result<Self> {
116        let rank = operand_dims.len();
117        if operand_strides.len() != rank || dest_dims.len() != rank || dest_strides.len() != rank {
118            return Err(StridedError::StrideLengthMismatch);
119        }
120        if starts.len() != rank {
121            return Err(StridedError::RankMismatch(starts.len(), rank));
122        }
123        if limits.len() != rank {
124            return Err(StridedError::RankMismatch(limits.len(), rank));
125        }
126        if slice_strides.len() != rank {
127            return Err(StridedError::RankMismatch(slice_strides.len(), rank));
128        }
129        checked_total_len(operand_dims)?;
130        checked_total_len(dest_dims)?;
131
132        let mut expected_dest_dims: AxisVec<usize> = AxisVec::with_capacity(rank);
133        let mut source_strides: AxisVec<isize> = AxisVec::with_capacity(rank);
134        let mut source_offset_delta = 0isize;
135        for axis in 0..rank {
136            let start = starts[axis];
137            let limit = limits[axis];
138            let stride = slice_strides[axis];
139            if start > limit || limit > operand_dims[axis] || stride == 0 {
140                return Err(StridedError::InvalidAxis { axis, rank });
141            }
142            let span = limit - start;
143            expected_dest_dims.push(span.div_ceil(stride));
144            source_strides.push(checked_stride_mul(operand_strides[axis], stride)?);
145            source_offset_delta =
146                checked_offset_add(source_offset_delta, operand_strides[axis], start)?;
147        }
148        if dest_dims != &expected_dest_dims[..] {
149            return Err(StridedError::ShapeMismatch(
150                dest_dims.to_vec(),
151                expected_dest_dims.to_vec(),
152            ));
153        }
154        let copy_plan = CopyPlan::compile(dest_dims, dest_strides, &source_strides)?;
155
156        Ok(Self {
157            operand_dims: operand_dims.into(),
158            operand_strides: operand_strides.into(),
159            dest_dims: dest_dims.into(),
160            dest_strides: dest_strides.into(),
161            source_strides,
162            source_offset_delta,
163            copy_plan,
164        })
165    }
166
167    /// Execute the prepared static slice traversal.
168    pub fn execute<T>(
169        &self,
170        dest: &mut RawStridedMut<'_, T>,
171        operand: &RawStridedRef<'_, T>,
172    ) -> Result<()>
173    where
174        T: Copy + MaybeSendSync,
175    {
176        self.check_call(dest, operand)?;
177        let source_offset = operand
178            .offset()
179            .checked_add(self.source_offset_delta)
180            .ok_or(StridedError::OffsetOverflow)?;
181        let source = unsafe {
182            RawStridedRef::new_unchecked(
183                operand.data(),
184                &self.dest_dims,
185                &self.source_strides,
186                source_offset,
187            )
188        };
189        self.copy_plan.execute(dest, &source)
190    }
191
192    /// Execute into storage whose reachable destination elements are uninitialized.
193    pub fn execute_uninit<T>(
194        &self,
195        dest: &mut RawStridedMut<'_, MaybeUninit<T>>,
196        operand: &RawStridedRef<'_, T>,
197    ) -> Result<()>
198    where
199        T: Copy + MaybeSendSync,
200    {
201        self.check_call(dest, operand)?;
202        let source_offset = operand
203            .offset()
204            .checked_add(self.source_offset_delta)
205            .ok_or(StridedError::OffsetOverflow)?;
206        let source = unsafe {
207            RawStridedRef::new_unchecked(
208                operand.data(),
209                &self.dest_dims,
210                &self.source_strides,
211                source_offset,
212            )
213        };
214        self.copy_plan.execute_uninit(dest, &source)
215    }
216
217    fn check_call<D, T>(
218        &self,
219        dest: &RawStridedMut<'_, D>,
220        operand: &RawStridedRef<'_, T>,
221    ) -> Result<()> {
222        if operand.dims() != &self.operand_dims[..]
223            || operand.strides() != &self.operand_strides[..]
224            || dest.dims() != &self.dest_dims[..]
225            || dest.strides() != &self.dest_strides[..]
226        {
227            return Err(StridedError::PlanLayoutMismatch);
228        }
229        Ok(())
230    }
231}
232
233impl PadPlan {
234    /// Compile a pad plan for one operand layout and destination layout.
235    #[allow(clippy::too_many_arguments)]
236    pub fn compile(
237        operand_dims: &[usize],
238        operand_strides: &[isize],
239        dest_dims: &[usize],
240        dest_strides: &[isize],
241        edge_padding_low: &[i64],
242        edge_padding_high: &[i64],
243        interior_padding: &[i64],
244    ) -> Result<Self> {
245        let rank = operand_dims.len();
246        if operand_strides.len() != rank || dest_dims.len() != rank || dest_strides.len() != rank {
247            return Err(StridedError::StrideLengthMismatch);
248        }
249        if edge_padding_low.len() != rank {
250            return Err(StridedError::RankMismatch(edge_padding_low.len(), rank));
251        }
252        if edge_padding_high.len() != rank {
253            return Err(StridedError::RankMismatch(edge_padding_high.len(), rank));
254        }
255        if interior_padding.len() != rank {
256            return Err(StridedError::RankMismatch(interior_padding.len(), rank));
257        }
258
259        let operand_total = checked_total_len(operand_dims)?;
260        let dest_total = checked_total_len(dest_dims)?;
261        if !crate::layout_check::is_injective_layout(dest_dims, dest_strides) {
262            return Err(StridedError::NonInjectiveOutputLayout);
263        }
264
265        let mut expected_dest_dims: AxisVec<usize> = AxisVec::with_capacity(rank);
266        let mut interior_step: AxisVec<i64> = AxisVec::with_capacity(rank);
267        for axis in 0..rank {
268            if interior_padding[axis] < 0 {
269                return Err(StridedError::InvalidAxis { axis, rank });
270            }
271            let step = interior_padding[axis]
272                .checked_add(1)
273                .ok_or(StridedError::OffsetOverflow)?;
274            interior_step.push(step);
275            expected_dest_dims.push(checked_pad_output_dim(
276                operand_dims[axis],
277                edge_padding_low[axis],
278                edge_padding_high[axis],
279                step,
280                axis,
281                rank,
282            )?);
283        }
284        if dest_dims != &expected_dest_dims[..] {
285            return Err(StridedError::ShapeMismatch(
286                dest_dims.to_vec(),
287                expected_dest_dims.to_vec(),
288            ));
289        }
290        let contiguous_dest_fill = is_dense_col_major(dest_dims, dest_strides);
291        let contiguous_axis0_run = compile_contiguous_pad_axis0_run(
292            operand_dims,
293            operand_strides,
294            dest_dims,
295            dest_strides,
296            edge_padding_low,
297            &interior_step,
298        );
299        let generic_fill = compile_pad_fill_cursor(dest_dims, dest_strides)?;
300        let generic_copy = compile_pad_copy_cursor(
301            operand_dims,
302            operand_strides,
303            dest_dims,
304            dest_strides,
305            edge_padding_low,
306            &interior_step,
307        )?;
308
309        Ok(Self {
310            operand_dims: operand_dims.into(),
311            operand_strides: operand_strides.into(),
312            dest_dims: dest_dims.into(),
313            dest_strides: dest_strides.into(),
314            edge_padding_low: edge_padding_low.into(),
315            interior_step,
316            operand_total,
317            dest_total,
318            contiguous_dest_fill,
319            contiguous_axis0_run,
320            generic_fill,
321            generic_copy,
322        })
323    }
324
325    /// Execute the prepared pad traversal.
326    pub fn execute<T>(
327        &self,
328        dest: &mut RawStridedMut<'_, T>,
329        operand: &RawStridedRef<'_, T>,
330        fill: T,
331    ) -> Result<()>
332    where
333        T: Copy + MaybeSendSync,
334    {
335        self.check_call(dest, operand)?;
336        self.fill_dest(dest, fill)?;
337
338        if self.operand_total == 0 {
339            return Ok(());
340        }
341        if let Some(run) = self.contiguous_axis0_run {
342            return self.copy_operand_axis0_runs(dest, operand, run);
343        }
344        self.copy_operand(dest, operand)
345    }
346
347    /// Execute pad into storage whose reachable destination elements are uninitialized.
348    pub fn execute_uninit<T>(
349        &self,
350        dest: &mut RawStridedMut<'_, MaybeUninit<T>>,
351        operand: &RawStridedRef<'_, T>,
352        fill: T,
353    ) -> Result<()>
354    where
355        T: Copy + MaybeSendSync,
356    {
357        self.check_call(dest, operand)?;
358        self.fill_dest(dest, MaybeUninit::new(fill))?;
359
360        if self.operand_total == 0 {
361            return Ok(());
362        }
363        if let Some(run) = self.contiguous_axis0_run {
364            return self.copy_operand_axis0_runs_uninit(dest, operand, run);
365        }
366        self.copy_operand_uninit(dest, operand)
367    }
368
369    fn fill_dest<T>(&self, dest: &mut RawStridedMut<'_, T>, fill: T) -> Result<()>
370    where
371        T: Copy + MaybeSendSync,
372    {
373        if self.dest_total == 0 {
374            return Ok(());
375        }
376        if self.contiguous_dest_fill {
377            let dest_offset =
378                usize::try_from(dest.offset()).map_err(|_| StridedError::OffsetOverflow)?;
379            let dest_end = dest_offset
380                .checked_add(self.dest_total)
381                .ok_or(StridedError::OffsetOverflow)?;
382            let dest_data = dest.data_mut();
383            let dest_slice = dest_data
384                .get_mut(dest_offset..dest_end)
385                .ok_or(StridedError::OffsetOverflow)?;
386            dest_slice.fill(fill);
387            return Ok(());
388        }
389        #[cfg(feature = "parallel")]
390        {
391            let nthreads = crate::threading::parallel_threads_for_len(self.dest_total);
392            if nthreads > 1 {
393                return self.fill_dest_parallel(dest, fill, nthreads);
394            }
395        }
396        self.fill_dest_serial(dest, fill)
397    }
398
399    fn copy_operand_axis0_runs<T>(
400        &self,
401        dest: &mut RawStridedMut<'_, T>,
402        operand: &RawStridedRef<'_, T>,
403        run: ContiguousPadAxis0Run,
404    ) -> Result<()>
405    where
406        T: Copy,
407    {
408        if run.len == 0 {
409            return Ok(());
410        }
411        let outer_dims = &self.operand_dims[1..];
412        let outer_total = checked_total_len(outer_dims)?;
413        let mut outer_idx_storage = CoordScratch::new(outer_dims.len());
414        let outer_idx = outer_idx_storage.as_mut_slice();
415        let operand_ptr = operand.data().as_ptr();
416        let dest_ptr = dest.data_mut().as_mut_ptr();
417
418        for _ in 0..outer_total {
419            let mut operand_offset =
420                checked_offset_add(operand.offset(), self.operand_strides[0], run.operand_start)?;
421            let mut dest_offset =
422                checked_offset_add(dest.offset(), self.dest_strides[0], run.dest_start)?;
423            let mut in_bounds = true;
424            for (outer_axis, &coord) in outer_idx.iter().enumerate() {
425                let axis = outer_axis + 1;
426                let out_pos = i128::from(self.edge_padding_low[axis])
427                    + coord as i128 * i128::from(self.interior_step[axis]);
428                if out_pos < 0 || out_pos >= self.dest_dims[axis] as i128 {
429                    in_bounds = false;
430                    break;
431                }
432                operand_offset =
433                    checked_offset_add(operand_offset, self.operand_strides[axis], coord)?;
434                dest_offset =
435                    checked_offset_add(dest_offset, self.dest_strides[axis], out_pos as usize)?;
436            }
437            if in_bounds {
438                unsafe {
439                    // SAFETY: compile requires unit axis-0 strides, the clipped
440                    // run is in bounds, and the destination layout is injective.
441                    // RawStridedMut's exclusive borrow cannot overlap the
442                    // RawStridedRef's shared borrow in safe Rust.
443                    core::ptr::copy_nonoverlapping(
444                        operand_ptr.offset(operand_offset),
445                        dest_ptr.offset(dest_offset),
446                        run.len,
447                    );
448                }
449            }
450            advance_col_major_index(outer_idx, outer_dims);
451        }
452        Ok(())
453    }
454
455    fn copy_operand_axis0_runs_uninit<T>(
456        &self,
457        dest: &mut RawStridedMut<'_, MaybeUninit<T>>,
458        operand: &RawStridedRef<'_, T>,
459        run: ContiguousPadAxis0Run,
460    ) -> Result<()>
461    where
462        T: Copy,
463    {
464        if run.len == 0 {
465            return Ok(());
466        }
467        let outer_dims = &self.operand_dims[1..];
468        let outer_total = checked_total_len(outer_dims)?;
469        let mut outer_idx_storage = CoordScratch::new(outer_dims.len());
470        let outer_idx = outer_idx_storage.as_mut_slice();
471        let operand_ptr = operand.data().as_ptr();
472        let dest_ptr = dest.data_mut().as_mut_ptr();
473
474        for _ in 0..outer_total {
475            let mut operand_offset =
476                checked_offset_add(operand.offset(), self.operand_strides[0], run.operand_start)?;
477            let mut dest_offset =
478                checked_offset_add(dest.offset(), self.dest_strides[0], run.dest_start)?;
479            let mut in_bounds = true;
480            for (outer_axis, &coord) in outer_idx.iter().enumerate() {
481                let axis = outer_axis + 1;
482                let out_pos = i128::from(self.edge_padding_low[axis])
483                    + coord as i128 * i128::from(self.interior_step[axis]);
484                if out_pos < 0 || out_pos >= self.dest_dims[axis] as i128 {
485                    in_bounds = false;
486                    break;
487                }
488                operand_offset =
489                    checked_offset_add(operand_offset, self.operand_strides[axis], coord)?;
490                dest_offset =
491                    checked_offset_add(dest_offset, self.dest_strides[axis], out_pos as usize)?;
492            }
493            if in_bounds {
494                unsafe {
495                    core::ptr::copy_nonoverlapping(
496                        operand_ptr.offset(operand_offset),
497                        dest_ptr.offset(dest_offset).cast::<T>(),
498                        run.len,
499                    );
500                }
501            }
502            advance_col_major_index(outer_idx, outer_dims);
503        }
504        Ok(())
505    }
506
507    #[cfg(test)]
508    fn contiguous_axis0_run(&self) -> Option<(usize, usize, usize)> {
509        self.contiguous_axis0_run
510            .map(|run| (run.operand_start, run.dest_start, run.len))
511    }
512
513    #[cfg(test)]
514    fn has_contiguous_dest_fill(&self) -> bool {
515        self.contiguous_dest_fill
516    }
517
518    fn fill_dest_serial<T>(&self, dest: &mut RawStridedMut<'_, T>, fill: T) -> Result<()>
519    where
520        T: Copy,
521    {
522        let dest_ptr = dest.data_mut().as_mut_ptr();
523        let mut cursor = PadFillState::new(dest.offset(), &self.dest_dims, &self.generic_fill);
524        for _ in 0..self.dest_total {
525            unsafe {
526                // INVARIANT: compile-time checked fill spans, validated raw
527                // descriptors, and exact plan-layout equality prove this is a
528                // reachable destination offset.
529                *dest_ptr.offset(cursor.offset) = fill;
530            }
531            cursor.advance(&self.dest_dims, &self.generic_fill);
532        }
533        Ok(())
534    }
535
536    #[cfg(feature = "parallel")]
537    fn fill_dest_parallel<T>(
538        &self,
539        dest: &mut RawStridedMut<'_, T>,
540        fill: T,
541        nthreads: usize,
542    ) -> Result<()>
543    where
544        T: Copy + MaybeSendSync,
545    {
546        let dest_offset_base = dest.offset();
547        let dest_ptr = crate::threading::SendPtr(dest.data_mut().as_mut_ptr());
548        crate::threading::parallel_map_reduce(
549            0..self.dest_total,
550            nthreads,
551            &|range| {
552                let mut cursor = PadFillState::decode(
553                    range.start,
554                    dest_offset_base,
555                    &self.dest_dims,
556                    &self.generic_fill,
557                )?;
558                let dest_ptr = dest_ptr.as_ptr();
559                for _ in range {
560                    unsafe {
561                        // INVARIANT: compile-time checked fill spans, validated
562                        // raw descriptors, exact layout equality, and disjoint
563                        // positive logical partition mapping prove this offset
564                        // is reachable and uniquely owned by this range.
565                        *dest_ptr.offset(cursor.offset) = fill;
566                    }
567                    cursor.advance(&self.dest_dims, &self.generic_fill);
568                }
569                Ok(())
570            },
571            &|left, right| left.and(right),
572        )
573    }
574
575    fn copy_operand<T>(
576        &self,
577        dest: &mut RawStridedMut<'_, T>,
578        operand: &RawStridedRef<'_, T>,
579    ) -> Result<()>
580    where
581        T: Copy + MaybeSendSync,
582    {
583        #[cfg(feature = "parallel")]
584        {
585            let nthreads = crate::threading::parallel_threads_for_len(self.generic_copy.total);
586            if nthreads > 1 {
587                return self.copy_operand_parallel(dest, operand, nthreads);
588            }
589        }
590        self.copy_operand_serial(dest, operand)
591    }
592
593    fn copy_operand_uninit<T>(
594        &self,
595        dest: &mut RawStridedMut<'_, MaybeUninit<T>>,
596        operand: &RawStridedRef<'_, T>,
597    ) -> Result<()>
598    where
599        T: Copy + MaybeSendSync,
600    {
601        #[cfg(feature = "parallel")]
602        {
603            let nthreads = crate::threading::parallel_threads_for_len(self.generic_copy.total);
604            if nthreads > 1 {
605                return self.copy_operand_uninit_parallel(dest, operand, nthreads);
606            }
607        }
608        self.copy_operand_uninit_serial(dest, operand)
609    }
610
611    fn copy_operand_uninit_serial<T>(
612        &self,
613        dest: &mut RawStridedMut<'_, MaybeUninit<T>>,
614        operand: &RawStridedRef<'_, T>,
615    ) -> Result<()>
616    where
617        T: Copy,
618    {
619        if self.generic_copy.total == 0 {
620            return Ok(());
621        }
622        let operand_ptr = operand.data().as_ptr();
623        let dest_ptr = dest.data_mut().as_mut_ptr();
624        let source_base = operand
625            .offset()
626            .checked_add(self.generic_copy.source_base_delta)
627            .ok_or(StridedError::OffsetOverflow)?;
628        let dest_base = dest
629            .offset()
630            .checked_add(self.generic_copy.dest_base_delta)
631            .ok_or(StridedError::OffsetOverflow)?;
632        let mut cursor = PadCopyState::new(source_base, dest_base, &self.generic_copy);
633        for _ in 0..self.generic_copy.total {
634            unsafe {
635                // INVARIANT: compile-time checked copy spans/base deltas,
636                // validated raw descriptors, and exact plan-layout equality
637                // prove both offsets reachable; positive steps make destinations
638                // injective, so no copy aliases another cursor position.
639                (*dest_ptr.offset(cursor.dest_offset))
640                    .write(*operand_ptr.offset(cursor.source_offset));
641            }
642            cursor.advance(&self.generic_copy);
643        }
644        Ok(())
645    }
646
647    #[cfg(feature = "parallel")]
648    fn copy_operand_uninit_parallel<T>(
649        &self,
650        dest: &mut RawStridedMut<'_, MaybeUninit<T>>,
651        operand: &RawStridedRef<'_, T>,
652        nthreads: usize,
653    ) -> Result<()>
654    where
655        T: Copy + MaybeSendSync,
656    {
657        let operand_base = operand
658            .offset()
659            .checked_add(self.generic_copy.source_base_delta)
660            .ok_or(StridedError::OffsetOverflow)?;
661        let operand_ptr = crate::threading::SendPtr(operand.data().as_ptr() as *mut T);
662        let dest_base = dest
663            .offset()
664            .checked_add(self.generic_copy.dest_base_delta)
665            .ok_or(StridedError::OffsetOverflow)?;
666        let dest_ptr = crate::threading::SendPtr(dest.data_mut().as_mut_ptr());
667        crate::threading::parallel_map_reduce(
668            0..self.generic_copy.total,
669            nthreads,
670            &|range| {
671                let mut cursor =
672                    PadCopyState::decode(range.start, operand_base, dest_base, &self.generic_copy)?;
673                let operand_ptr = operand_ptr.as_const();
674                let dest_ptr = dest_ptr.as_ptr();
675
676                for _ in range {
677                    unsafe {
678                        // INVARIANT: compile-time checked copy spans/base
679                        // deltas, validated descriptors, exact layout equality,
680                        // and positive injective mapping prove disjoint,
681                        // reachable source/destination offsets per range.
682                        (*dest_ptr.offset(cursor.dest_offset))
683                            .write(*operand_ptr.offset(cursor.source_offset));
684                    }
685                    cursor.advance(&self.generic_copy);
686                }
687                Ok(())
688            },
689            &|left, right| left.and(right),
690        )
691    }
692
693    fn copy_operand_serial<T>(
694        &self,
695        dest: &mut RawStridedMut<'_, T>,
696        operand: &RawStridedRef<'_, T>,
697    ) -> Result<()>
698    where
699        T: Copy,
700    {
701        if self.generic_copy.total == 0 {
702            return Ok(());
703        }
704        let operand_ptr = operand.data().as_ptr();
705        let dest_ptr = dest.data_mut().as_mut_ptr();
706        let source_base = operand
707            .offset()
708            .checked_add(self.generic_copy.source_base_delta)
709            .ok_or(StridedError::OffsetOverflow)?;
710        let dest_base = dest
711            .offset()
712            .checked_add(self.generic_copy.dest_base_delta)
713            .ok_or(StridedError::OffsetOverflow)?;
714        let mut cursor = PadCopyState::new(source_base, dest_base, &self.generic_copy);
715        for _ in 0..self.generic_copy.total {
716            unsafe {
717                // INVARIANT: compile-time checked copy spans/base deltas,
718                // validated raw descriptors, and exact plan-layout equality
719                // prove both offsets reachable; positive steps make destinations
720                // injective, so no copy aliases another cursor position.
721                *dest_ptr.offset(cursor.dest_offset) = *operand_ptr.offset(cursor.source_offset);
722            }
723            cursor.advance(&self.generic_copy);
724        }
725        Ok(())
726    }
727
728    #[cfg(feature = "parallel")]
729    fn copy_operand_parallel<T>(
730        &self,
731        dest: &mut RawStridedMut<'_, T>,
732        operand: &RawStridedRef<'_, T>,
733        nthreads: usize,
734    ) -> Result<()>
735    where
736        T: Copy + MaybeSendSync,
737    {
738        let operand_base = operand
739            .offset()
740            .checked_add(self.generic_copy.source_base_delta)
741            .ok_or(StridedError::OffsetOverflow)?;
742        let operand_ptr = crate::threading::SendPtr(operand.data().as_ptr() as *mut T);
743        let dest_base = dest
744            .offset()
745            .checked_add(self.generic_copy.dest_base_delta)
746            .ok_or(StridedError::OffsetOverflow)?;
747        let dest_ptr = crate::threading::SendPtr(dest.data_mut().as_mut_ptr());
748        crate::threading::parallel_map_reduce(
749            0..self.generic_copy.total,
750            nthreads,
751            &|range| {
752                let mut cursor =
753                    PadCopyState::decode(range.start, operand_base, dest_base, &self.generic_copy)?;
754                let operand_ptr = operand_ptr.as_const();
755                let dest_ptr = dest_ptr.as_ptr();
756
757                for _ in range {
758                    unsafe {
759                        // INVARIANT: compile-time checked copy spans/base
760                        // deltas, validated descriptors, exact layout equality,
761                        // and positive injective mapping prove disjoint,
762                        // reachable source/destination offsets per range.
763                        *dest_ptr.offset(cursor.dest_offset) =
764                            *operand_ptr.offset(cursor.source_offset);
765                    }
766                    cursor.advance(&self.generic_copy);
767                }
768                Ok(())
769            },
770            &|left, right| left.and(right),
771        )
772    }
773
774    fn check_call<D, T>(
775        &self,
776        dest: &RawStridedMut<'_, D>,
777        operand: &RawStridedRef<'_, T>,
778    ) -> Result<()> {
779        if operand.dims() != &self.operand_dims[..]
780            || operand.strides() != &self.operand_strides[..]
781            || dest.dims() != &self.dest_dims[..]
782            || dest.strides() != &self.dest_strides[..]
783        {
784            return Err(StridedError::PlanLayoutMismatch);
785        }
786        Ok(())
787    }
788}
789
790impl ConcatenatePlan {
791    /// Compile a multi-input concatenate plan for fixed input and destination layouts.
792    pub fn compile(
793        input_dims: &[&[usize]],
794        input_strides: &[&[isize]],
795        dest_dims: &[usize],
796        dest_strides: &[isize],
797        axis: usize,
798    ) -> Result<Self> {
799        if input_dims.is_empty() {
800            return Err(StridedError::UnsupportedArity {
801                arity: 0,
802                max: usize::MAX,
803            });
804        }
805        if input_dims.len() != input_strides.len() {
806            return Err(StridedError::RankMismatch(
807                input_strides.len(),
808                input_dims.len(),
809            ));
810        }
811
812        let rank = input_dims[0].len();
813        if dest_dims.len() != rank || dest_strides.len() != rank {
814            return Err(StridedError::StrideLengthMismatch);
815        }
816        if axis >= rank {
817            return Err(StridedError::InvalidAxis { axis, rank });
818        }
819        checked_total_len(dest_dims)?;
820        if !crate::layout_check::is_injective_layout(dest_dims, dest_strides) {
821            return Err(StridedError::NonInjectiveOutputLayout);
822        }
823
824        let mut expected_dest_dims: AxisVec<usize> = input_dims[0].into();
825        expected_dest_dims[axis] = 0;
826        let mut stored_input_dims = Vec::with_capacity(input_dims.len());
827        let mut stored_input_strides = Vec::with_capacity(input_dims.len());
828        let mut dest_offset_deltas = Vec::with_capacity(input_dims.len());
829        let mut copy_plans = Vec::with_capacity(input_dims.len());
830        let mut axis_base = 0usize;
831
832        for (dims, strides) in input_dims.iter().zip(input_strides.iter()) {
833            if dims.len() != rank {
834                return Err(StridedError::RankMismatch(dims.len(), rank));
835            }
836            if strides.len() != rank {
837                return Err(StridedError::StrideLengthMismatch);
838            }
839            checked_total_len(dims)?;
840            for dim in 0..rank {
841                if dim == axis {
842                    expected_dest_dims[axis] = expected_dest_dims[axis]
843                        .checked_add(dims[axis])
844                        .ok_or(StridedError::OffsetOverflow)?;
845                } else if dims[dim] != input_dims[0][dim] {
846                    return Err(StridedError::ShapeMismatch(
847                        dims.to_vec(),
848                        input_dims[0].to_vec(),
849                    ));
850                }
851            }
852            dest_offset_deltas.push(checked_offset_add(0, dest_strides[axis], axis_base)?);
853            axis_base = axis_base
854                .checked_add(dims[axis])
855                .ok_or(StridedError::OffsetOverflow)?;
856            copy_plans.push(CopyPlan::compile(dims, dest_strides, strides)?);
857            stored_input_dims.push((*dims).into());
858            stored_input_strides.push((*strides).into());
859        }
860
861        if dest_dims != &expected_dest_dims[..] {
862            return Err(StridedError::ShapeMismatch(
863                dest_dims.to_vec(),
864                expected_dest_dims.to_vec(),
865            ));
866        }
867
868        Ok(Self {
869            input_dims: stored_input_dims,
870            input_strides: stored_input_strides,
871            dest_dims: dest_dims.into(),
872            dest_strides: dest_strides.into(),
873            dest_offset_deltas,
874            copy_plans,
875        })
876    }
877
878    /// Execute the prepared concatenate traversal.
879    pub fn execute<T>(
880        &self,
881        dest: &mut RawStridedMut<'_, T>,
882        inputs: &[RawStridedRef<'_, T>],
883    ) -> Result<()>
884    where
885        T: Copy + MaybeSendSync,
886    {
887        self.check_dest_layout(dest)?;
888        if inputs.len() != self.input_dims.len() {
889            return Err(StridedError::RankMismatch(
890                inputs.len(),
891                self.input_dims.len(),
892            ));
893        }
894        for (position, input) in inputs.iter().enumerate() {
895            self.check_input_layout(position, input)?;
896            self.segment_offset(position, dest.offset())?;
897        }
898        for (position, input) in inputs.iter().enumerate() {
899            self.execute_segment(position, dest, input)?;
900        }
901        Ok(())
902    }
903
904    /// Execute concatenate into storage whose reachable destination elements are uninitialized.
905    pub fn execute_uninit<T>(
906        &self,
907        dest: &mut RawStridedMut<'_, MaybeUninit<T>>,
908        inputs: &[RawStridedRef<'_, T>],
909    ) -> Result<()>
910    where
911        T: Copy + MaybeSendSync,
912    {
913        self.check_dest_layout(dest)?;
914        if inputs.len() != self.input_dims.len() {
915            return Err(StridedError::RankMismatch(
916                inputs.len(),
917                self.input_dims.len(),
918            ));
919        }
920        for (position, input) in inputs.iter().enumerate() {
921            self.check_input_layout(position, input)?;
922            self.segment_offset(position, dest.offset())?;
923        }
924        for (position, input) in inputs.iter().enumerate() {
925            self.execute_segment_uninit(position, dest, input)?;
926        }
927        Ok(())
928    }
929
930    pub(crate) fn check_dest_layout<T>(&self, dest: &RawStridedMut<'_, T>) -> Result<()> {
931        if dest.dims() != &self.dest_dims[..] || dest.strides() != &self.dest_strides[..] {
932            return Err(StridedError::PlanLayoutMismatch);
933        }
934        Ok(())
935    }
936
937    pub(crate) fn check_input_layout<T>(
938        &self,
939        position: usize,
940        input: &RawStridedRef<'_, T>,
941    ) -> Result<()> {
942        if position >= self.input_dims.len()
943            || input.dims() != &self.input_dims[position][..]
944            || input.strides() != &self.input_strides[position][..]
945        {
946            return Err(StridedError::PlanLayoutMismatch);
947        }
948        Ok(())
949    }
950
951    pub(crate) fn input_count(&self) -> usize {
952        self.input_dims.len()
953    }
954
955    pub(crate) fn segment_offset(&self, position: usize, dest_offset: isize) -> Result<isize> {
956        dest_offset
957            .checked_add(self.dest_offset_deltas[position])
958            .ok_or(StridedError::OffsetOverflow)
959    }
960
961    pub(crate) fn execute_segment<T>(
962        &self,
963        position: usize,
964        dest: &mut RawStridedMut<'_, T>,
965        input: &RawStridedRef<'_, T>,
966    ) -> Result<()>
967    where
968        T: Copy + MaybeSendSync,
969    {
970        let segment_offset = self.segment_offset(position, dest.offset())?;
971        let dest_data = dest.data_mut();
972        let mut segment = unsafe {
973            RawStridedMut::new_unchecked(
974                dest_data,
975                &self.input_dims[position],
976                &self.dest_strides,
977                segment_offset,
978            )
979        };
980        self.copy_plans[position].execute(&mut segment, input)
981    }
982
983    pub(crate) fn execute_segment_uninit<T>(
984        &self,
985        position: usize,
986        dest: &mut RawStridedMut<'_, MaybeUninit<T>>,
987        input: &RawStridedRef<'_, T>,
988    ) -> Result<()>
989    where
990        T: Copy + MaybeSendSync,
991    {
992        let segment_offset = self.segment_offset(position, dest.offset())?;
993        let dest_data = dest.data_mut();
994        let mut segment = unsafe {
995            RawStridedMut::new_unchecked(
996                dest_data,
997                &self.input_dims[position],
998                &self.dest_strides,
999                segment_offset,
1000            )
1001        };
1002        self.copy_plans[position].execute_uninit(&mut segment, input)
1003    }
1004}
1005
1006impl ReversePlan {
1007    /// Compile a reverse plan for one operand layout, destination layout, and axis set.
1008    pub fn compile(
1009        operand_dims: &[usize],
1010        operand_strides: &[isize],
1011        dest_strides: &[isize],
1012        axes: &[usize],
1013    ) -> Result<Self> {
1014        let rank = operand_dims.len();
1015        if operand_strides.len() != rank || dest_strides.len() != rank {
1016            return Err(StridedError::StrideLengthMismatch);
1017        }
1018        checked_total_len(operand_dims)?;
1019
1020        let mut reverse_axis: AxisVec<bool> = (0..rank).map(|_| false).collect();
1021        for &axis in axes {
1022            if axis >= rank {
1023                return Err(StridedError::InvalidAxis { axis, rank });
1024            }
1025            reverse_axis[axis] = true;
1026        }
1027
1028        let mut source_strides: AxisVec<isize> = AxisVec::with_capacity(rank);
1029        let mut source_offset_delta = 0isize;
1030        for axis in 0..rank {
1031            if reverse_axis[axis] {
1032                source_strides.push(
1033                    operand_strides[axis]
1034                        .checked_neg()
1035                        .ok_or(StridedError::OffsetOverflow)?,
1036                );
1037                if operand_dims[axis] > 0 {
1038                    source_offset_delta = checked_offset_add(
1039                        source_offset_delta,
1040                        operand_strides[axis],
1041                        operand_dims[axis] - 1,
1042                    )?;
1043                }
1044            } else {
1045                source_strides.push(operand_strides[axis]);
1046            }
1047        }
1048        let copy_plan = CopyPlan::compile(operand_dims, dest_strides, &source_strides)?;
1049
1050        Ok(Self {
1051            operand_dims: operand_dims.into(),
1052            operand_strides: operand_strides.into(),
1053            dest_strides: dest_strides.into(),
1054            source_strides,
1055            source_offset_delta,
1056            copy_plan,
1057        })
1058    }
1059
1060    /// Execute the prepared reverse traversal.
1061    pub fn execute<T>(
1062        &self,
1063        dest: &mut RawStridedMut<'_, T>,
1064        operand: &RawStridedRef<'_, T>,
1065    ) -> Result<()>
1066    where
1067        T: Copy + MaybeSendSync,
1068    {
1069        self.check_call(dest, operand)?;
1070        let source_offset = operand
1071            .offset()
1072            .checked_add(self.source_offset_delta)
1073            .ok_or(StridedError::OffsetOverflow)?;
1074        let source = unsafe {
1075            RawStridedRef::new_unchecked(
1076                operand.data(),
1077                &self.operand_dims,
1078                &self.source_strides,
1079                source_offset,
1080            )
1081        };
1082        self.copy_plan.execute(dest, &source)
1083    }
1084
1085    /// Execute reverse into storage whose reachable destination elements are uninitialized.
1086    pub fn execute_uninit<T>(
1087        &self,
1088        dest: &mut RawStridedMut<'_, MaybeUninit<T>>,
1089        operand: &RawStridedRef<'_, T>,
1090    ) -> Result<()>
1091    where
1092        T: Copy + MaybeSendSync,
1093    {
1094        self.check_call(dest, operand)?;
1095        let source_offset = operand
1096            .offset()
1097            .checked_add(self.source_offset_delta)
1098            .ok_or(StridedError::OffsetOverflow)?;
1099        let source = unsafe {
1100            RawStridedRef::new_unchecked(
1101                operand.data(),
1102                &self.operand_dims,
1103                &self.source_strides,
1104                source_offset,
1105            )
1106        };
1107        self.copy_plan.execute_uninit(dest, &source)
1108    }
1109
1110    fn check_call<D, T>(
1111        &self,
1112        dest: &RawStridedMut<'_, D>,
1113        operand: &RawStridedRef<'_, T>,
1114    ) -> Result<()> {
1115        if operand.dims() != &self.operand_dims[..]
1116            || operand.strides() != &self.operand_strides[..]
1117            || dest.dims() != &self.operand_dims[..]
1118            || dest.strides() != &self.dest_strides[..]
1119        {
1120            return Err(StridedError::PlanLayoutMismatch);
1121        }
1122        Ok(())
1123    }
1124}
1125
1126fn checked_total_len(dims: &[usize]) -> Result<usize> {
1127    if dims.is_empty() {
1128        return Ok(1);
1129    }
1130    dims.iter()
1131        .try_fold(1usize, |acc, &dim| acc.checked_mul(dim))
1132        .ok_or(StridedError::OffsetOverflow)
1133}
1134
1135fn checked_stride_mul(stride: isize, factor: usize) -> Result<isize> {
1136    let factor = isize::try_from(factor).map_err(|_| StridedError::OffsetOverflow)?;
1137    stride
1138        .checked_mul(factor)
1139        .ok_or(StridedError::OffsetOverflow)
1140}
1141
1142fn checked_pad_output_dim(
1143    input_extent: usize,
1144    edge_low: i64,
1145    edge_high: i64,
1146    interior_step: i64,
1147    axis: usize,
1148    rank: usize,
1149) -> Result<usize> {
1150    let base = if input_extent == 0 {
1151        0i128
1152    } else {
1153        (input_extent as i128 - 1)
1154            .checked_mul(i128::from(interior_step))
1155            .and_then(|value| value.checked_add(1))
1156            .ok_or(StridedError::OffsetOverflow)?
1157    };
1158    let dim = i128::from(edge_low)
1159        .checked_add(i128::from(edge_high))
1160        .and_then(|value| value.checked_add(base))
1161        .ok_or(StridedError::OffsetOverflow)?;
1162    usize::try_from(dim).map_err(|_| StridedError::InvalidAxis { axis, rank })
1163}
1164
1165fn compile_pad_fill_cursor(dims: &[usize], strides: &[isize]) -> Result<PadFillCursor> {
1166    let mut steps = AxisVec::with_capacity(dims.len());
1167    let mut resets = AxisVec::with_capacity(dims.len());
1168    for (&dim, &stride) in dims.iter().zip(strides) {
1169        steps.push(stride);
1170        resets.push(checked_cursor_reset(stride, dim)?);
1171    }
1172    check_offset_span(dims, &steps)?;
1173    Ok(PadFillCursor { steps, resets })
1174}
1175
1176fn compile_pad_copy_cursor(
1177    operand_dims: &[usize],
1178    operand_strides: &[isize],
1179    dest_dims: &[usize],
1180    dest_strides: &[isize],
1181    edge_padding_low: &[i64],
1182    interior_step: &[i64],
1183) -> Result<PadCopyCursor> {
1184    let mut shape = AxisVec::with_capacity(operand_dims.len());
1185    let mut source_steps = AxisVec::with_capacity(operand_dims.len());
1186    let mut source_resets = AxisVec::with_capacity(operand_dims.len());
1187    let mut dest_steps = AxisVec::with_capacity(operand_dims.len());
1188    let mut dest_resets = AxisVec::with_capacity(operand_dims.len());
1189    let mut source_base_delta = 0isize;
1190    let mut dest_base_delta = 0isize;
1191    let mut copy_empty = false;
1192
1193    for axis in 0..operand_dims.len() {
1194        let (start, end) = checked_pad_valid_interval(
1195            operand_dims[axis],
1196            dest_dims[axis],
1197            edge_padding_low[axis],
1198            interior_step[axis],
1199        )?;
1200        let extent = end - start;
1201        shape.push(extent);
1202        if !copy_empty && extent != 0 {
1203            source_base_delta =
1204                checked_offset_add(source_base_delta, operand_strides[axis], start)?;
1205            let output_start = i128::from(edge_padding_low[axis])
1206                .checked_add(
1207                    i128::try_from(start)
1208                        .map_err(|_| StridedError::OffsetOverflow)?
1209                        .checked_mul(i128::from(interior_step[axis]))
1210                        .ok_or(StridedError::OffsetOverflow)?,
1211                )
1212                .ok_or(StridedError::OffsetOverflow)?;
1213            let output_start =
1214                usize::try_from(output_start).map_err(|_| StridedError::OffsetOverflow)?;
1215            dest_base_delta =
1216                checked_offset_add(dest_base_delta, dest_strides[axis], output_start)?;
1217        }
1218        copy_empty |= extent == 0;
1219
1220        let source_step = operand_strides[axis];
1221        let dest_step = checked_stride_mul_i64(dest_strides[axis], interior_step[axis])?;
1222        source_steps.push(source_step);
1223        source_resets.push(checked_cursor_reset(source_step, extent)?);
1224        dest_steps.push(dest_step);
1225        dest_resets.push(checked_cursor_reset(dest_step, extent)?);
1226    }
1227
1228    check_offset_span(&shape, &source_steps)?;
1229    check_offset_span(&shape, &dest_steps)?;
1230    let total = if shape.iter().any(|&extent| extent == 0) {
1231        0
1232    } else {
1233        checked_total_len(&shape)?
1234    };
1235
1236    Ok(PadCopyCursor {
1237        shape,
1238        source_base_delta,
1239        dest_base_delta,
1240        source_steps,
1241        source_resets,
1242        dest_steps,
1243        dest_resets,
1244        total,
1245    })
1246}
1247
1248fn checked_pad_valid_interval(
1249    input_extent: usize,
1250    dest_extent: usize,
1251    edge_low: i64,
1252    step: i64,
1253) -> Result<(usize, usize)> {
1254    if input_extent == 0 || dest_extent == 0 {
1255        return Ok((0, 0));
1256    }
1257    let step = i128::from(step);
1258    let lower = ceil_div_positive(-i128::from(edge_low), step)?;
1259    let dest_last = i128::try_from(dest_extent)
1260        .map_err(|_| StridedError::OffsetOverflow)?
1261        .checked_sub(1)
1262        .ok_or(StridedError::OffsetOverflow)?;
1263    let upper = floor_div_positive(
1264        dest_last
1265            .checked_sub(i128::from(edge_low))
1266            .ok_or(StridedError::OffsetOverflow)?,
1267        step,
1268    )?
1269    .checked_add(1)
1270    .ok_or(StridedError::OffsetOverflow)?;
1271    let input_extent = i128::try_from(input_extent).map_err(|_| StridedError::OffsetOverflow)?;
1272    let lower = lower.clamp(0, input_extent);
1273    let upper = upper.clamp(0, input_extent);
1274    if lower >= upper {
1275        return Ok((0, 0));
1276    }
1277    Ok((
1278        usize::try_from(lower).map_err(|_| StridedError::OffsetOverflow)?,
1279        usize::try_from(upper).map_err(|_| StridedError::OffsetOverflow)?,
1280    ))
1281}
1282
1283fn ceil_div_positive(numerator: i128, denominator: i128) -> Result<i128> {
1284    if denominator <= 0 {
1285        return Err(StridedError::OffsetOverflow);
1286    }
1287    let quotient = numerator.div_euclid(denominator);
1288    let remainder = numerator.rem_euclid(denominator);
1289    quotient
1290        .checked_add(i128::from(remainder != 0))
1291        .ok_or(StridedError::OffsetOverflow)
1292}
1293
1294fn floor_div_positive(numerator: i128, denominator: i128) -> Result<i128> {
1295    if denominator <= 0 {
1296        return Err(StridedError::OffsetOverflow);
1297    }
1298    Ok(numerator.div_euclid(denominator))
1299}
1300
1301fn checked_stride_mul_i64(stride: isize, factor: i64) -> Result<isize> {
1302    let factor = isize::try_from(factor).map_err(|_| StridedError::OffsetOverflow)?;
1303    stride
1304        .checked_mul(factor)
1305        .ok_or(StridedError::OffsetOverflow)
1306}
1307
1308fn checked_cursor_reset(step: isize, extent: usize) -> Result<isize> {
1309    if extent == 0 {
1310        return Ok(0);
1311    }
1312    let last = isize::try_from(extent - 1).map_err(|_| StridedError::OffsetOverflow)?;
1313    step.checked_mul(last)
1314        .and_then(isize::checked_neg)
1315        .ok_or(StridedError::OffsetOverflow)
1316}
1317
1318fn check_offset_span(shape: &[usize], strides: &[isize]) -> Result<()> {
1319    let mut min = 0isize;
1320    let mut max = 0isize;
1321    for (&extent, &stride) in shape.iter().zip(strides) {
1322        if extent <= 1 {
1323            continue;
1324        }
1325        let last = isize::try_from(extent - 1).map_err(|_| StridedError::OffsetOverflow)?;
1326        let delta = stride
1327            .checked_mul(last)
1328            .ok_or(StridedError::OffsetOverflow)?;
1329        if delta < 0 {
1330            min = min.checked_add(delta).ok_or(StridedError::OffsetOverflow)?;
1331        } else {
1332            max = max.checked_add(delta).ok_or(StridedError::OffsetOverflow)?;
1333        }
1334    }
1335    let _ = (min, max);
1336    Ok(())
1337}
1338
1339fn is_dense_col_major(dims: &[usize], strides: &[isize]) -> bool {
1340    let mut expected = 1isize;
1341    for (&dim, &stride) in dims.iter().zip(strides.iter()) {
1342        if stride != expected {
1343            return false;
1344        }
1345        let Ok(dim) = isize::try_from(dim) else {
1346            return false;
1347        };
1348        let Some(next) = expected.checked_mul(dim) else {
1349            return false;
1350        };
1351        expected = next;
1352    }
1353    true
1354}
1355
1356fn compile_contiguous_pad_axis0_run(
1357    operand_dims: &[usize],
1358    operand_strides: &[isize],
1359    dest_dims: &[usize],
1360    dest_strides: &[isize],
1361    edge_padding_low: &[i64],
1362    interior_step: &[i64],
1363) -> Option<ContiguousPadAxis0Run> {
1364    if operand_dims.is_empty()
1365        || operand_strides[0] != 1
1366        || dest_strides[0] != 1
1367        || interior_step[0] != 1
1368    {
1369        return None;
1370    }
1371
1372    let operand_extent = operand_dims[0] as i128;
1373    let dest_extent = dest_dims[0] as i128;
1374    let edge_low = i128::from(edge_padding_low[0]);
1375    let operand_start = (-edge_low).clamp(0, operand_extent);
1376    let dest_start = edge_low.clamp(0, dest_extent);
1377    let len = (operand_extent - operand_start).min(dest_extent - dest_start);
1378    Some(ContiguousPadAxis0Run {
1379        operand_start: usize::try_from(operand_start).ok()?,
1380        dest_start: usize::try_from(dest_start).ok()?,
1381        len: usize::try_from(len).ok()?,
1382    })
1383}
1384
1385fn checked_offset_add(base: isize, stride: isize, coord: usize) -> Result<isize> {
1386    let coord = isize::try_from(coord).map_err(|_| StridedError::OffsetOverflow)?;
1387    let scaled = stride
1388        .checked_mul(coord)
1389        .ok_or(StridedError::OffsetOverflow)?;
1390    base.checked_add(scaled).ok_or(StridedError::OffsetOverflow)
1391}
1392
1393fn advance_col_major_index(index: &mut [usize], shape: &[usize]) {
1394    for axis in 0..index.len() {
1395        index[axis] += 1;
1396        if index[axis] < shape[axis] {
1397            return;
1398        }
1399        index[axis] = 0;
1400    }
1401}
1402
1403#[cfg(feature = "parallel")]
1404fn fill_col_major_index(mut linear: usize, shape: &[usize], out: &mut [usize]) {
1405    for (axis, coord) in out.iter_mut().enumerate() {
1406        let dim = shape[axis];
1407        *coord = linear % dim;
1408        linear /= dim;
1409    }
1410}
1411
1412struct PadFillState {
1413    coords: CoordScratch,
1414    offset: isize,
1415}
1416
1417impl PadFillState {
1418    fn new(base: isize, shape: &[usize], _cursor: &PadFillCursor) -> Self {
1419        Self {
1420            coords: CoordScratch::new(shape.len()),
1421            offset: base,
1422        }
1423    }
1424
1425    #[cfg(feature = "parallel")]
1426    fn decode(linear: usize, base: isize, shape: &[usize], cursor: &PadFillCursor) -> Result<Self> {
1427        let mut state = Self::new(base, shape, cursor);
1428        fill_col_major_index(linear, shape, state.coords.as_mut_slice());
1429        for (&coord, &step) in state.coords.as_mut_slice().iter().zip(&cursor.steps) {
1430            state.offset = checked_offset_add(state.offset, step, coord)?;
1431        }
1432        Ok(state)
1433    }
1434
1435    #[inline]
1436    fn advance(&mut self, shape: &[usize], cursor: &PadFillCursor) {
1437        // INVARIANT: precomputed checked offset arithmetic, the validated
1438        // descriptor, and exact layout equality prove each incremental offset
1439        // remains reachable.
1440        for axis in 0..shape.len() {
1441            let next = self.coords.as_mut_slice()[axis] + 1;
1442            if next < shape[axis] {
1443                self.coords.as_mut_slice()[axis] = next;
1444                self.offset += cursor.steps[axis];
1445                return;
1446            }
1447            self.coords.as_mut_slice()[axis] = 0;
1448            self.offset += cursor.resets[axis];
1449        }
1450    }
1451}
1452
1453struct PadCopyState {
1454    coords: CoordScratch,
1455    source_offset: isize,
1456    dest_offset: isize,
1457}
1458
1459impl PadCopyState {
1460    fn new(source_base: isize, dest_base: isize, cursor: &PadCopyCursor) -> Self {
1461        Self {
1462            coords: CoordScratch::new(cursor.shape.len()),
1463            source_offset: source_base,
1464            dest_offset: dest_base,
1465        }
1466    }
1467
1468    #[cfg(feature = "parallel")]
1469    fn decode(
1470        linear: usize,
1471        source_base: isize,
1472        dest_base: isize,
1473        cursor: &PadCopyCursor,
1474    ) -> Result<Self> {
1475        let mut state = Self::new(source_base, dest_base, cursor);
1476        fill_col_major_index(linear, &cursor.shape, state.coords.as_mut_slice());
1477        for axis in 0..cursor.shape.len() {
1478            let coord = state.coords.as_mut_slice()[axis];
1479            state.source_offset =
1480                checked_offset_add(state.source_offset, cursor.source_steps[axis], coord)?;
1481            state.dest_offset =
1482                checked_offset_add(state.dest_offset, cursor.dest_steps[axis], coord)?;
1483        }
1484        Ok(state)
1485    }
1486
1487    #[inline]
1488    fn advance(&mut self, cursor: &PadCopyCursor) {
1489        // INVARIANT: compile-time checked steps/resets/offset arithmetic,
1490        // validated descriptors, and exact plan-layout equality prove every update below
1491        // stays in the reachable source/destination domains.
1492        for axis in 0..cursor.shape.len() {
1493            let next = self.coords.as_mut_slice()[axis] + 1;
1494            if next < cursor.shape[axis] {
1495                self.coords.as_mut_slice()[axis] = next;
1496                self.source_offset += cursor.source_steps[axis];
1497                self.dest_offset += cursor.dest_steps[axis];
1498                return;
1499            }
1500            self.coords.as_mut_slice()[axis] = 0;
1501            self.source_offset += cursor.source_resets[axis];
1502            self.dest_offset += cursor.dest_resets[axis];
1503        }
1504    }
1505}
1506
1507struct CoordScratch {
1508    inline: [usize; crate::RAW_FUSED_RANK_LIMIT],
1509    heap: Option<Vec<usize>>,
1510    len: usize,
1511}
1512
1513impl CoordScratch {
1514    fn new(len: usize) -> Self {
1515        if len <= crate::RAW_FUSED_RANK_LIMIT {
1516            Self {
1517                inline: [0; crate::RAW_FUSED_RANK_LIMIT],
1518                heap: None,
1519                len,
1520            }
1521        } else {
1522            Self {
1523                inline: [0; crate::RAW_FUSED_RANK_LIMIT],
1524                heap: Some(vec![0; len]),
1525                len,
1526            }
1527        }
1528    }
1529
1530    fn as_mut_slice(&mut self) -> &mut [usize] {
1531        match &mut self.heap {
1532            Some(heap) => heap,
1533            None => &mut self.inline[..self.len],
1534        }
1535    }
1536}
1537
1538#[cfg(test)]
1539#[path = "static_indexing_plan/tests/tests.rs"]
1540mod tests;