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    /// Logical start of each segment in the concatenated index space
102    /// (segment totals prefix sums); the last entry is the destination total.
103    #[cfg_attr(not(feature = "parallel"), allow(dead_code))]
104    segment_starts: Vec<usize>,
105    copy_plans: Vec<CopyPlan>,
106}
107
108impl SlicePlan {
109    /// Compile a static slice plan for one operand layout and destination layout.
110    #[allow(clippy::too_many_arguments)]
111    pub fn compile(
112        operand_dims: &[usize],
113        operand_strides: &[isize],
114        dest_dims: &[usize],
115        dest_strides: &[isize],
116        starts: &[usize],
117        limits: &[usize],
118        slice_strides: &[usize],
119    ) -> Result<Self> {
120        let rank = operand_dims.len();
121        if operand_strides.len() != rank || dest_dims.len() != rank || dest_strides.len() != rank {
122            return Err(StridedError::StrideLengthMismatch);
123        }
124        if starts.len() != rank {
125            return Err(StridedError::RankMismatch(starts.len(), rank));
126        }
127        if limits.len() != rank {
128            return Err(StridedError::RankMismatch(limits.len(), rank));
129        }
130        if slice_strides.len() != rank {
131            return Err(StridedError::RankMismatch(slice_strides.len(), rank));
132        }
133        checked_total_len(operand_dims)?;
134        checked_total_len(dest_dims)?;
135
136        let mut expected_dest_dims: AxisVec<usize> = AxisVec::with_capacity(rank);
137        let mut source_strides: AxisVec<isize> = AxisVec::with_capacity(rank);
138        let mut source_offset_delta = 0isize;
139        for axis in 0..rank {
140            let start = starts[axis];
141            let limit = limits[axis];
142            let stride = slice_strides[axis];
143            if start > limit || limit > operand_dims[axis] || stride == 0 {
144                return Err(StridedError::InvalidAxis { axis, rank });
145            }
146            let span = limit - start;
147            expected_dest_dims.push(span.div_ceil(stride));
148            source_strides.push(checked_stride_mul(operand_strides[axis], stride)?);
149            source_offset_delta =
150                checked_offset_add(source_offset_delta, operand_strides[axis], start)?;
151        }
152        if dest_dims != &expected_dest_dims[..] {
153            return Err(StridedError::ShapeMismatch(
154                dest_dims.to_vec(),
155                expected_dest_dims.to_vec(),
156            ));
157        }
158        let copy_plan = CopyPlan::compile(dest_dims, dest_strides, &source_strides)?;
159
160        Ok(Self {
161            operand_dims: operand_dims.into(),
162            operand_strides: operand_strides.into(),
163            dest_dims: dest_dims.into(),
164            dest_strides: dest_strides.into(),
165            source_strides,
166            source_offset_delta,
167            copy_plan,
168        })
169    }
170
171    /// Execute the prepared static slice traversal.
172    pub fn execute<T>(
173        &self,
174        dest: &mut RawStridedMut<'_, T>,
175        operand: &RawStridedRef<'_, T>,
176    ) -> Result<()>
177    where
178        T: Copy + MaybeSendSync,
179    {
180        self.check_call(dest, operand)?;
181        let source_offset = operand
182            .offset()
183            .checked_add(self.source_offset_delta)
184            .ok_or(StridedError::OffsetOverflow)?;
185        let source = unsafe {
186            RawStridedRef::new_unchecked(
187                operand.data(),
188                &self.dest_dims,
189                &self.source_strides,
190                source_offset,
191            )
192        };
193        self.copy_plan.execute(dest, &source)
194    }
195
196    /// Execute into storage whose reachable destination elements are uninitialized.
197    pub fn execute_uninit<T>(
198        &self,
199        dest: &mut RawStridedMut<'_, MaybeUninit<T>>,
200        operand: &RawStridedRef<'_, T>,
201    ) -> Result<()>
202    where
203        T: Copy + MaybeSendSync,
204    {
205        self.check_call(dest, operand)?;
206        let source_offset = operand
207            .offset()
208            .checked_add(self.source_offset_delta)
209            .ok_or(StridedError::OffsetOverflow)?;
210        let source = unsafe {
211            RawStridedRef::new_unchecked(
212                operand.data(),
213                &self.dest_dims,
214                &self.source_strides,
215                source_offset,
216            )
217        };
218        self.copy_plan.execute_uninit(dest, &source)
219    }
220
221    fn check_call<D, T>(
222        &self,
223        dest: &RawStridedMut<'_, D>,
224        operand: &RawStridedRef<'_, T>,
225    ) -> Result<()> {
226        if operand.dims() != &self.operand_dims[..]
227            || operand.strides() != &self.operand_strides[..]
228            || dest.dims() != &self.dest_dims[..]
229            || dest.strides() != &self.dest_strides[..]
230        {
231            return Err(StridedError::PlanLayoutMismatch);
232        }
233        Ok(())
234    }
235}
236
237impl PadPlan {
238    /// Compile a pad plan for one operand layout and destination layout.
239    #[allow(clippy::too_many_arguments)]
240    pub fn compile(
241        operand_dims: &[usize],
242        operand_strides: &[isize],
243        dest_dims: &[usize],
244        dest_strides: &[isize],
245        edge_padding_low: &[i64],
246        edge_padding_high: &[i64],
247        interior_padding: &[i64],
248    ) -> Result<Self> {
249        let rank = operand_dims.len();
250        if operand_strides.len() != rank || dest_dims.len() != rank || dest_strides.len() != rank {
251            return Err(StridedError::StrideLengthMismatch);
252        }
253        if edge_padding_low.len() != rank {
254            return Err(StridedError::RankMismatch(edge_padding_low.len(), rank));
255        }
256        if edge_padding_high.len() != rank {
257            return Err(StridedError::RankMismatch(edge_padding_high.len(), rank));
258        }
259        if interior_padding.len() != rank {
260            return Err(StridedError::RankMismatch(interior_padding.len(), rank));
261        }
262
263        let operand_total = checked_total_len(operand_dims)?;
264        let dest_total = checked_total_len(dest_dims)?;
265        if !crate::layout_check::is_injective_layout(dest_dims, dest_strides) {
266            return Err(StridedError::NonInjectiveOutputLayout);
267        }
268
269        let mut expected_dest_dims: AxisVec<usize> = AxisVec::with_capacity(rank);
270        let mut interior_step: AxisVec<i64> = AxisVec::with_capacity(rank);
271        for axis in 0..rank {
272            if interior_padding[axis] < 0 {
273                return Err(StridedError::InvalidAxis { axis, rank });
274            }
275            let step = interior_padding[axis]
276                .checked_add(1)
277                .ok_or(StridedError::OffsetOverflow)?;
278            interior_step.push(step);
279            expected_dest_dims.push(checked_pad_output_dim(
280                operand_dims[axis],
281                edge_padding_low[axis],
282                edge_padding_high[axis],
283                step,
284                axis,
285                rank,
286            )?);
287        }
288        if dest_dims != &expected_dest_dims[..] {
289            return Err(StridedError::ShapeMismatch(
290                dest_dims.to_vec(),
291                expected_dest_dims.to_vec(),
292            ));
293        }
294        let contiguous_dest_fill = is_dense_col_major(dest_dims, dest_strides);
295        let contiguous_axis0_run = compile_contiguous_pad_axis0_run(
296            operand_dims,
297            operand_strides,
298            dest_dims,
299            dest_strides,
300            edge_padding_low,
301            &interior_step,
302        );
303        let generic_fill = compile_pad_fill_cursor(dest_dims, dest_strides)?;
304        let generic_copy = compile_pad_copy_cursor(
305            operand_dims,
306            operand_strides,
307            dest_dims,
308            dest_strides,
309            edge_padding_low,
310            &interior_step,
311        )?;
312
313        Ok(Self {
314            operand_dims: operand_dims.into(),
315            operand_strides: operand_strides.into(),
316            dest_dims: dest_dims.into(),
317            dest_strides: dest_strides.into(),
318            edge_padding_low: edge_padding_low.into(),
319            interior_step,
320            operand_total,
321            dest_total,
322            contiguous_dest_fill,
323            contiguous_axis0_run,
324            generic_fill,
325            generic_copy,
326        })
327    }
328
329    /// Execute the prepared pad traversal.
330    pub fn execute<T>(
331        &self,
332        dest: &mut RawStridedMut<'_, T>,
333        operand: &RawStridedRef<'_, T>,
334        fill: T,
335    ) -> Result<()>
336    where
337        T: Copy + MaybeSendSync,
338    {
339        self.check_call(dest, operand)?;
340        self.fill_dest(dest, fill)?;
341
342        if self.operand_total == 0 {
343            return Ok(());
344        }
345        if let Some(run) = self.contiguous_axis0_run {
346            return self.copy_operand_axis0_runs(dest, operand, run);
347        }
348        self.copy_operand(dest, operand)
349    }
350
351    /// Execute pad into storage whose reachable destination elements are uninitialized.
352    pub fn execute_uninit<T>(
353        &self,
354        dest: &mut RawStridedMut<'_, MaybeUninit<T>>,
355        operand: &RawStridedRef<'_, T>,
356        fill: T,
357    ) -> Result<()>
358    where
359        T: Copy + MaybeSendSync,
360    {
361        self.check_call(dest, operand)?;
362        self.fill_dest(dest, MaybeUninit::new(fill))?;
363
364        if self.operand_total == 0 {
365            return Ok(());
366        }
367        if let Some(run) = self.contiguous_axis0_run {
368            return self.copy_operand_axis0_runs_uninit(dest, operand, run);
369        }
370        self.copy_operand_uninit(dest, operand)
371    }
372
373    fn fill_dest<T>(&self, dest: &mut RawStridedMut<'_, T>, fill: T) -> Result<()>
374    where
375        T: Copy + MaybeSendSync,
376    {
377        if self.dest_total == 0 {
378            return Ok(());
379        }
380        if self.contiguous_dest_fill {
381            let dest_offset =
382                usize::try_from(dest.offset()).map_err(|_| StridedError::OffsetOverflow)?;
383            let dest_end = dest_offset
384                .checked_add(self.dest_total)
385                .ok_or(StridedError::OffsetOverflow)?;
386            let dest_data = dest.data_mut();
387            let dest_slice = dest_data
388                .get_mut(dest_offset..dest_end)
389                .ok_or(StridedError::OffsetOverflow)?;
390            crate::threading::fill_contiguous(dest_slice, fill);
391            return Ok(());
392        }
393        #[cfg(feature = "parallel")]
394        {
395            let nthreads = crate::threading::parallel_threads_for_len(self.dest_total);
396            if nthreads > 1 {
397                return self.fill_dest_parallel(dest, fill, nthreads);
398            }
399        }
400        self.fill_dest_serial(dest, fill)
401    }
402
403    fn copy_operand_axis0_runs<T>(
404        &self,
405        dest: &mut RawStridedMut<'_, T>,
406        operand: &RawStridedRef<'_, T>,
407        run: ContiguousPadAxis0Run,
408    ) -> Result<()>
409    where
410        T: Copy + MaybeSendSync,
411    {
412        let dest_offset = dest.offset();
413        let dest_ptr = dest.data_mut().as_mut_ptr();
414        // SAFETY: `check_call` proved `dest` carries the compiled destination
415        // layout, so its validated data covers every reachable destination
416        // slot, and the exclusive borrow is the only access during the call.
417        unsafe { self.copy_axis0_runs_raw(dest_ptr, dest_offset, operand, run) }
418    }
419
420    fn copy_operand_axis0_runs_uninit<T>(
421        &self,
422        dest: &mut RawStridedMut<'_, MaybeUninit<T>>,
423        operand: &RawStridedRef<'_, T>,
424        run: ContiguousPadAxis0Run,
425    ) -> Result<()>
426    where
427        T: Copy + MaybeSendSync,
428    {
429        let dest_offset = dest.offset();
430        let dest_ptr = dest.data_mut().as_mut_ptr().cast::<T>();
431        // SAFETY: as in `copy_operand_axis0_runs`; `MaybeUninit<T>` has the
432        // layout of `T`, and only writes go through the cast pointer.
433        unsafe { self.copy_axis0_runs_raw(dest_ptr, dest_offset, operand, run) }
434    }
435
436    /// Copy every in-bounds axis-0 run of the operand into the destination.
437    ///
438    /// With at least as many runs as workers (above the repository
439    /// threshold) the runs are split across workers; otherwise runs are
440    /// copied in order and each long run is chunked across workers by
441    /// [`crate::threading::copy_contiguous`].
442    ///
443    /// # Safety
444    ///
445    /// `dest_ptr` must be the data pointer of a destination whose layout is
446    /// exactly the compiled destination layout from `dest_offset`, valid for
447    /// writes to every reachable slot, with no concurrent access.
448    unsafe fn copy_axis0_runs_raw<T>(
449        &self,
450        dest_ptr: *mut T,
451        dest_offset: isize,
452        operand: &RawStridedRef<'_, T>,
453        run: ContiguousPadAxis0Run,
454    ) -> Result<()>
455    where
456        T: Copy + MaybeSendSync,
457    {
458        if run.len == 0 {
459            return Ok(());
460        }
461        let outer_dims = &self.operand_dims[1..];
462        let outer_total = checked_total_len(outer_dims)?;
463        #[cfg(feature = "parallel")]
464        {
465            let copied = outer_total.saturating_mul(run.len);
466            let nthreads = crate::threading::parallel_threads_for_len(copied);
467            if nthreads > 1 && outer_total >= nthreads {
468                let dest_ptr = crate::threading::SendPtr(dest_ptr);
469                return crate::threading::parallel_map_reduce(
470                    0..outer_total,
471                    nthreads,
472                    &|range| {
473                        let mut outer_idx_storage = CoordScratch::new(outer_dims.len());
474                        let outer_idx = outer_idx_storage.as_mut_slice();
475                        fill_col_major_index(range.start, outer_dims, outer_idx);
476                        // SAFETY: the caller's contract covers every run, and
477                        // disjoint outer ranges write disjoint runs because the
478                        // compiled destination layout is injective.
479                        unsafe {
480                            self.copy_axis0_run_range(
481                                dest_ptr.as_ptr(),
482                                dest_offset,
483                                operand,
484                                run,
485                                outer_idx,
486                                range.len(),
487                            )
488                        }
489                    },
490                    &|left, right| left.and(right),
491                );
492            }
493        }
494        let mut outer_idx_storage = CoordScratch::new(outer_dims.len());
495        let outer_idx = outer_idx_storage.as_mut_slice();
496        // SAFETY: forwarded from this function's contract.
497        unsafe {
498            self.copy_axis0_run_range(dest_ptr, dest_offset, operand, run, outer_idx, outer_total)
499        }
500    }
501
502    /// Copy `count` consecutive axis-0 runs starting at outer coordinate
503    /// `outer_idx` (column-major over axes `1..`), advancing it in place.
504    ///
505    /// # Safety
506    ///
507    /// As for [`Self::copy_axis0_runs_raw`], and no other thread may write the
508    /// destination runs visited by this call.
509    unsafe fn copy_axis0_run_range<T>(
510        &self,
511        dest_ptr: *mut T,
512        dest_offset: isize,
513        operand: &RawStridedRef<'_, T>,
514        run: ContiguousPadAxis0Run,
515        outer_idx: &mut [usize],
516        count: usize,
517    ) -> Result<()>
518    where
519        T: Copy + MaybeSendSync,
520    {
521        let outer_dims = &self.operand_dims[1..];
522        let operand_ptr = operand.data().as_ptr();
523        for _ in 0..count {
524            let mut operand_offset =
525                checked_offset_add(operand.offset(), self.operand_strides[0], run.operand_start)?;
526            let mut dest_run_offset =
527                checked_offset_add(dest_offset, self.dest_strides[0], run.dest_start)?;
528            let mut in_bounds = true;
529            for (outer_axis, &coord) in outer_idx.iter().enumerate() {
530                let axis = outer_axis + 1;
531                let out_pos = i128::from(self.edge_padding_low[axis])
532                    + coord as i128 * i128::from(self.interior_step[axis]);
533                if out_pos < 0 || out_pos >= self.dest_dims[axis] as i128 {
534                    in_bounds = false;
535                    break;
536                }
537                operand_offset =
538                    checked_offset_add(operand_offset, self.operand_strides[axis], coord)?;
539                dest_run_offset =
540                    checked_offset_add(dest_run_offset, self.dest_strides[axis], out_pos as usize)?;
541            }
542            if in_bounds {
543                // SAFETY: compile requires unit axis-0 strides, the clipped
544                // run is in bounds of the validated operand and destination,
545                // and the destination layout is injective. RawStridedMut's
546                // exclusive borrow cannot overlap the RawStridedRef's shared
547                // borrow in safe Rust.
548                unsafe {
549                    crate::threading::copy_contiguous(
550                        operand_ptr.offset(operand_offset),
551                        dest_ptr.offset(dest_run_offset),
552                        run.len,
553                    );
554                }
555            }
556            advance_col_major_index(outer_idx, outer_dims);
557        }
558        Ok(())
559    }
560
561    #[cfg(test)]
562    fn contiguous_axis0_run(&self) -> Option<(usize, usize, usize)> {
563        self.contiguous_axis0_run
564            .map(|run| (run.operand_start, run.dest_start, run.len))
565    }
566
567    #[cfg(test)]
568    fn has_contiguous_dest_fill(&self) -> bool {
569        self.contiguous_dest_fill
570    }
571
572    fn fill_dest_serial<T>(&self, dest: &mut RawStridedMut<'_, T>, fill: T) -> Result<()>
573    where
574        T: Copy,
575    {
576        let dest_ptr = dest.data_mut().as_mut_ptr();
577        let mut cursor = PadFillState::new(dest.offset(), &self.dest_dims, &self.generic_fill);
578        for _ in 0..self.dest_total {
579            unsafe {
580                // INVARIANT: compile-time checked fill spans, validated raw
581                // descriptors, and exact plan-layout equality prove this is a
582                // reachable destination offset.
583                *dest_ptr.offset(cursor.offset) = fill;
584            }
585            cursor.advance(&self.dest_dims, &self.generic_fill);
586        }
587        Ok(())
588    }
589
590    #[cfg(feature = "parallel")]
591    fn fill_dest_parallel<T>(
592        &self,
593        dest: &mut RawStridedMut<'_, T>,
594        fill: T,
595        nthreads: usize,
596    ) -> Result<()>
597    where
598        T: Copy + MaybeSendSync,
599    {
600        let dest_offset_base = dest.offset();
601        let dest_ptr = crate::threading::SendPtr(dest.data_mut().as_mut_ptr());
602        crate::threading::parallel_map_reduce(
603            0..self.dest_total,
604            nthreads,
605            &|range| {
606                let mut cursor = PadFillState::decode(
607                    range.start,
608                    dest_offset_base,
609                    &self.dest_dims,
610                    &self.generic_fill,
611                )?;
612                let dest_ptr = dest_ptr.as_ptr();
613                for _ in range {
614                    unsafe {
615                        // INVARIANT: compile-time checked fill spans, validated
616                        // raw descriptors, exact layout equality, and disjoint
617                        // positive logical partition mapping prove this offset
618                        // is reachable and uniquely owned by this range.
619                        *dest_ptr.offset(cursor.offset) = fill;
620                    }
621                    cursor.advance(&self.dest_dims, &self.generic_fill);
622                }
623                Ok(())
624            },
625            &|left, right| left.and(right),
626        )
627    }
628
629    fn copy_operand<T>(
630        &self,
631        dest: &mut RawStridedMut<'_, T>,
632        operand: &RawStridedRef<'_, T>,
633    ) -> Result<()>
634    where
635        T: Copy + MaybeSendSync,
636    {
637        #[cfg(feature = "parallel")]
638        {
639            let nthreads = crate::threading::parallel_threads_for_len(self.generic_copy.total);
640            if nthreads > 1 {
641                return self.copy_operand_parallel(dest, operand, nthreads);
642            }
643        }
644        self.copy_operand_serial(dest, operand)
645    }
646
647    fn copy_operand_uninit<T>(
648        &self,
649        dest: &mut RawStridedMut<'_, MaybeUninit<T>>,
650        operand: &RawStridedRef<'_, T>,
651    ) -> Result<()>
652    where
653        T: Copy + MaybeSendSync,
654    {
655        #[cfg(feature = "parallel")]
656        {
657            let nthreads = crate::threading::parallel_threads_for_len(self.generic_copy.total);
658            if nthreads > 1 {
659                return self.copy_operand_uninit_parallel(dest, operand, nthreads);
660            }
661        }
662        self.copy_operand_uninit_serial(dest, operand)
663    }
664
665    fn copy_operand_uninit_serial<T>(
666        &self,
667        dest: &mut RawStridedMut<'_, MaybeUninit<T>>,
668        operand: &RawStridedRef<'_, T>,
669    ) -> Result<()>
670    where
671        T: Copy,
672    {
673        if self.generic_copy.total == 0 {
674            return Ok(());
675        }
676        let operand_ptr = operand.data().as_ptr();
677        let dest_ptr = dest.data_mut().as_mut_ptr();
678        let source_base = operand
679            .offset()
680            .checked_add(self.generic_copy.source_base_delta)
681            .ok_or(StridedError::OffsetOverflow)?;
682        let dest_base = dest
683            .offset()
684            .checked_add(self.generic_copy.dest_base_delta)
685            .ok_or(StridedError::OffsetOverflow)?;
686        let mut cursor = PadCopyState::new(source_base, dest_base, &self.generic_copy);
687        for _ in 0..self.generic_copy.total {
688            unsafe {
689                // INVARIANT: compile-time checked copy spans/base deltas,
690                // validated raw descriptors, and exact plan-layout equality
691                // prove both offsets reachable; positive steps make destinations
692                // injective, so no copy aliases another cursor position.
693                (*dest_ptr.offset(cursor.dest_offset))
694                    .write(*operand_ptr.offset(cursor.source_offset));
695            }
696            cursor.advance(&self.generic_copy);
697        }
698        Ok(())
699    }
700
701    #[cfg(feature = "parallel")]
702    fn copy_operand_uninit_parallel<T>(
703        &self,
704        dest: &mut RawStridedMut<'_, MaybeUninit<T>>,
705        operand: &RawStridedRef<'_, T>,
706        nthreads: usize,
707    ) -> Result<()>
708    where
709        T: Copy + MaybeSendSync,
710    {
711        let operand_base = operand
712            .offset()
713            .checked_add(self.generic_copy.source_base_delta)
714            .ok_or(StridedError::OffsetOverflow)?;
715        let operand_ptr = crate::threading::SendPtr(operand.data().as_ptr() as *mut T);
716        let dest_base = dest
717            .offset()
718            .checked_add(self.generic_copy.dest_base_delta)
719            .ok_or(StridedError::OffsetOverflow)?;
720        let dest_ptr = crate::threading::SendPtr(dest.data_mut().as_mut_ptr());
721        crate::threading::parallel_map_reduce(
722            0..self.generic_copy.total,
723            nthreads,
724            &|range| {
725                let mut cursor =
726                    PadCopyState::decode(range.start, operand_base, dest_base, &self.generic_copy)?;
727                let operand_ptr = operand_ptr.as_const();
728                let dest_ptr = dest_ptr.as_ptr();
729
730                for _ in range {
731                    unsafe {
732                        // INVARIANT: compile-time checked copy spans/base
733                        // deltas, validated descriptors, exact layout equality,
734                        // and positive injective mapping prove disjoint,
735                        // reachable source/destination offsets per range.
736                        (*dest_ptr.offset(cursor.dest_offset))
737                            .write(*operand_ptr.offset(cursor.source_offset));
738                    }
739                    cursor.advance(&self.generic_copy);
740                }
741                Ok(())
742            },
743            &|left, right| left.and(right),
744        )
745    }
746
747    fn copy_operand_serial<T>(
748        &self,
749        dest: &mut RawStridedMut<'_, T>,
750        operand: &RawStridedRef<'_, T>,
751    ) -> Result<()>
752    where
753        T: Copy,
754    {
755        if self.generic_copy.total == 0 {
756            return Ok(());
757        }
758        let operand_ptr = operand.data().as_ptr();
759        let dest_ptr = dest.data_mut().as_mut_ptr();
760        let source_base = operand
761            .offset()
762            .checked_add(self.generic_copy.source_base_delta)
763            .ok_or(StridedError::OffsetOverflow)?;
764        let dest_base = dest
765            .offset()
766            .checked_add(self.generic_copy.dest_base_delta)
767            .ok_or(StridedError::OffsetOverflow)?;
768        let mut cursor = PadCopyState::new(source_base, dest_base, &self.generic_copy);
769        for _ in 0..self.generic_copy.total {
770            unsafe {
771                // INVARIANT: compile-time checked copy spans/base deltas,
772                // validated raw descriptors, and exact plan-layout equality
773                // prove both offsets reachable; positive steps make destinations
774                // injective, so no copy aliases another cursor position.
775                *dest_ptr.offset(cursor.dest_offset) = *operand_ptr.offset(cursor.source_offset);
776            }
777            cursor.advance(&self.generic_copy);
778        }
779        Ok(())
780    }
781
782    #[cfg(feature = "parallel")]
783    fn copy_operand_parallel<T>(
784        &self,
785        dest: &mut RawStridedMut<'_, T>,
786        operand: &RawStridedRef<'_, T>,
787        nthreads: usize,
788    ) -> Result<()>
789    where
790        T: Copy + MaybeSendSync,
791    {
792        let operand_base = operand
793            .offset()
794            .checked_add(self.generic_copy.source_base_delta)
795            .ok_or(StridedError::OffsetOverflow)?;
796        let operand_ptr = crate::threading::SendPtr(operand.data().as_ptr() as *mut T);
797        let dest_base = dest
798            .offset()
799            .checked_add(self.generic_copy.dest_base_delta)
800            .ok_or(StridedError::OffsetOverflow)?;
801        let dest_ptr = crate::threading::SendPtr(dest.data_mut().as_mut_ptr());
802        crate::threading::parallel_map_reduce(
803            0..self.generic_copy.total,
804            nthreads,
805            &|range| {
806                let mut cursor =
807                    PadCopyState::decode(range.start, operand_base, dest_base, &self.generic_copy)?;
808                let operand_ptr = operand_ptr.as_const();
809                let dest_ptr = dest_ptr.as_ptr();
810
811                for _ in range {
812                    unsafe {
813                        // INVARIANT: compile-time checked copy spans/base
814                        // deltas, validated descriptors, exact layout equality,
815                        // and positive injective mapping prove disjoint,
816                        // reachable source/destination offsets per range.
817                        *dest_ptr.offset(cursor.dest_offset) =
818                            *operand_ptr.offset(cursor.source_offset);
819                    }
820                    cursor.advance(&self.generic_copy);
821                }
822                Ok(())
823            },
824            &|left, right| left.and(right),
825        )
826    }
827
828    fn check_call<D, T>(
829        &self,
830        dest: &RawStridedMut<'_, D>,
831        operand: &RawStridedRef<'_, T>,
832    ) -> Result<()> {
833        if operand.dims() != &self.operand_dims[..]
834            || operand.strides() != &self.operand_strides[..]
835            || dest.dims() != &self.dest_dims[..]
836            || dest.strides() != &self.dest_strides[..]
837        {
838            return Err(StridedError::PlanLayoutMismatch);
839        }
840        Ok(())
841    }
842}
843
844#[cfg(feature = "parallel")]
845struct ConcatSegment<'a, T> {
846    layout: &'a crate::raw_ops::FusedPairLayout,
847    dest_offset: isize,
848    src_ptr: crate::threading::SendPtr<T>,
849    src_offset: isize,
850}
851
852impl ConcatenatePlan {
853    /// Compile a multi-input concatenate plan for fixed input and destination layouts.
854    pub fn compile(
855        input_dims: &[&[usize]],
856        input_strides: &[&[isize]],
857        dest_dims: &[usize],
858        dest_strides: &[isize],
859        axis: usize,
860    ) -> Result<Self> {
861        if input_dims.is_empty() {
862            return Err(StridedError::UnsupportedArity {
863                arity: 0,
864                max: usize::MAX,
865            });
866        }
867        if input_dims.len() != input_strides.len() {
868            return Err(StridedError::RankMismatch(
869                input_strides.len(),
870                input_dims.len(),
871            ));
872        }
873
874        let rank = input_dims[0].len();
875        if dest_dims.len() != rank || dest_strides.len() != rank {
876            return Err(StridedError::StrideLengthMismatch);
877        }
878        if axis >= rank {
879            return Err(StridedError::InvalidAxis { axis, rank });
880        }
881        checked_total_len(dest_dims)?;
882        if !crate::layout_check::is_injective_layout(dest_dims, dest_strides) {
883            return Err(StridedError::NonInjectiveOutputLayout);
884        }
885
886        let mut expected_dest_dims: AxisVec<usize> = input_dims[0].into();
887        expected_dest_dims[axis] = 0;
888        let mut stored_input_dims = Vec::with_capacity(input_dims.len());
889        let mut stored_input_strides = Vec::with_capacity(input_dims.len());
890        let mut dest_offset_deltas = Vec::with_capacity(input_dims.len());
891        let mut copy_plans = Vec::with_capacity(input_dims.len());
892        let mut segment_starts = Vec::with_capacity(input_dims.len() + 1);
893        segment_starts.push(0usize);
894        let mut axis_base = 0usize;
895
896        for (dims, strides) in input_dims.iter().zip(input_strides.iter()) {
897            if dims.len() != rank {
898                return Err(StridedError::RankMismatch(dims.len(), rank));
899            }
900            if strides.len() != rank {
901                return Err(StridedError::StrideLengthMismatch);
902            }
903            let segment_total = checked_total_len(dims)?;
904            let segment_end = segment_starts[segment_starts.len() - 1]
905                .checked_add(segment_total)
906                .ok_or(StridedError::OffsetOverflow)?;
907            segment_starts.push(segment_end);
908            for dim in 0..rank {
909                if dim == axis {
910                    expected_dest_dims[axis] = expected_dest_dims[axis]
911                        .checked_add(dims[axis])
912                        .ok_or(StridedError::OffsetOverflow)?;
913                } else if dims[dim] != input_dims[0][dim] {
914                    return Err(StridedError::ShapeMismatch(
915                        dims.to_vec(),
916                        input_dims[0].to_vec(),
917                    ));
918                }
919            }
920            dest_offset_deltas.push(checked_offset_add(0, dest_strides[axis], axis_base)?);
921            axis_base = axis_base
922                .checked_add(dims[axis])
923                .ok_or(StridedError::OffsetOverflow)?;
924            copy_plans.push(CopyPlan::compile(dims, dest_strides, strides)?);
925            stored_input_dims.push((*dims).into());
926            stored_input_strides.push((*strides).into());
927        }
928
929        if dest_dims != &expected_dest_dims[..] {
930            return Err(StridedError::ShapeMismatch(
931                dest_dims.to_vec(),
932                expected_dest_dims.to_vec(),
933            ));
934        }
935
936        Ok(Self {
937            input_dims: stored_input_dims,
938            input_strides: stored_input_strides,
939            dest_dims: dest_dims.into(),
940            dest_strides: dest_strides.into(),
941            dest_offset_deltas,
942            segment_starts,
943            copy_plans,
944        })
945    }
946
947    /// Execute the prepared concatenate traversal.
948    pub fn execute<T>(
949        &self,
950        dest: &mut RawStridedMut<'_, T>,
951        inputs: &[RawStridedRef<'_, T>],
952    ) -> Result<()>
953    where
954        T: Copy + MaybeSendSync,
955    {
956        self.check_dest_layout(dest)?;
957        if inputs.len() != self.input_dims.len() {
958            return Err(StridedError::RankMismatch(
959                inputs.len(),
960                self.input_dims.len(),
961            ));
962        }
963        for (position, input) in inputs.iter().enumerate() {
964            self.check_input_layout(position, input)?;
965            self.segment_offset(position, dest.offset())?;
966        }
967        #[cfg(feature = "parallel")]
968        if self.try_execute_parallel(dest, inputs, |dst: &mut T, value| *dst = value)? {
969            return Ok(());
970        }
971        for (position, input) in inputs.iter().enumerate() {
972            self.execute_segment(position, dest, input)?;
973        }
974        Ok(())
975    }
976
977    /// Execute concatenate into storage whose reachable destination elements are uninitialized.
978    pub fn execute_uninit<T>(
979        &self,
980        dest: &mut RawStridedMut<'_, MaybeUninit<T>>,
981        inputs: &[RawStridedRef<'_, T>],
982    ) -> Result<()>
983    where
984        T: Copy + MaybeSendSync,
985    {
986        self.check_dest_layout(dest)?;
987        if inputs.len() != self.input_dims.len() {
988            return Err(StridedError::RankMismatch(
989                inputs.len(),
990                self.input_dims.len(),
991            ));
992        }
993        for (position, input) in inputs.iter().enumerate() {
994            self.check_input_layout(position, input)?;
995            self.segment_offset(position, dest.offset())?;
996        }
997        #[cfg(feature = "parallel")]
998        if self.try_execute_parallel(dest, inputs, |dst: &mut MaybeUninit<T>, value| {
999            dst.write(value);
1000        })? {
1001            return Ok(());
1002        }
1003        for (position, input) in inputs.iter().enumerate() {
1004            self.execute_segment_uninit(position, dest, input)?;
1005        }
1006        Ok(())
1007    }
1008
1009    /// Replay all segments as one logical index space split across workers.
1010    ///
1011    /// Splitting the concatenated space (instead of fanning out per segment)
1012    /// keeps many small segments parallel as a group and splits a large
1013    /// segment across workers, with one fork-join for the whole call. Returns
1014    /// `Ok(false)` without writing when the serial per-segment path must run:
1015    /// the destination is at most the repository threshold, the active policy
1016    /// allows one thread, or a segment exceeds the fused-rank limit.
1017    ///
1018    /// The caller must have validated every input layout and segment offset.
1019    #[cfg(feature = "parallel")]
1020    fn try_execute_parallel<D, T, Apply>(
1021        &self,
1022        dest: &mut RawStridedMut<'_, D>,
1023        inputs: &[RawStridedRef<'_, T>],
1024        apply: Apply,
1025    ) -> Result<bool>
1026    where
1027        D: Copy + MaybeSendSync,
1028        T: Copy + MaybeSendSync,
1029        Apply: Fn(&mut D, T) + MaybeSendSync,
1030    {
1031        let total = self.segment_starts[self.segment_starts.len() - 1];
1032        let nthreads = crate::threading::parallel_threads_for_len(total);
1033        if nthreads <= 1 {
1034            return Ok(false);
1035        }
1036        let mut segments = Vec::with_capacity(inputs.len());
1037        for (position, input) in inputs.iter().enumerate() {
1038            let Some(layout) = self.copy_plans[position].fused_layout() else {
1039                return Ok(false);
1040            };
1041            segments.push(ConcatSegment {
1042                layout,
1043                dest_offset: self.segment_offset(position, dest.offset())?,
1044                src_ptr: crate::threading::SendPtr(input.data().as_ptr() as *mut T),
1045                src_offset: input.offset(),
1046            });
1047        }
1048        let dest_ptr = crate::threading::SendPtr(dest.data_mut().as_mut_ptr());
1049        let starts = &self.segment_starts;
1050        let segments = &segments;
1051        crate::threading::parallel_for_each(0..total, nthreads, &|range| {
1052            // First segment whose end lies past the range start.
1053            let mut position = starts[1..].partition_point(|&end| end <= range.start);
1054            while position < segments.len() && starts[position] < range.end {
1055                let segment = &segments[position];
1056                let local_start = range.start.max(starts[position]) - starts[position];
1057                let local_end = range.end.min(starts[position + 1]) - starts[position];
1058                // SAFETY: `check_dest_layout`/`check_input_layout` proved the
1059                // views carry the compiled layouts, so each segment window
1060                // (input dims, destination strides, checked segment offset) is
1061                // a sub-layout of the validated destination and each input
1062                // layout is in-bounds of its data. The destination layout is
1063                // injective and the segments occupy disjoint ranges of the
1064                // concatenation axis, so disjoint logical ranges write
1065                // disjoint slots; inputs are shared borrows that cannot alias
1066                // the exclusive destination borrow.
1067                unsafe {
1068                    crate::raw_ops::apply_fused_range(
1069                        dest_ptr.as_ptr(),
1070                        segment.dest_offset,
1071                        segment.src_ptr.as_const(),
1072                        segment.src_offset,
1073                        segment.layout,
1074                        local_start,
1075                        local_end - local_start,
1076                        &apply,
1077                        &|value| value,
1078                    );
1079                }
1080                position += 1;
1081            }
1082        });
1083        Ok(true)
1084    }
1085
1086    pub(crate) fn check_dest_layout<T>(&self, dest: &RawStridedMut<'_, T>) -> Result<()> {
1087        if dest.dims() != &self.dest_dims[..] || dest.strides() != &self.dest_strides[..] {
1088            return Err(StridedError::PlanLayoutMismatch);
1089        }
1090        Ok(())
1091    }
1092
1093    pub(crate) fn check_input_layout<T>(
1094        &self,
1095        position: usize,
1096        input: &RawStridedRef<'_, T>,
1097    ) -> Result<()> {
1098        if position >= self.input_dims.len()
1099            || input.dims() != &self.input_dims[position][..]
1100            || input.strides() != &self.input_strides[position][..]
1101        {
1102            return Err(StridedError::PlanLayoutMismatch);
1103        }
1104        Ok(())
1105    }
1106
1107    pub(crate) fn input_count(&self) -> usize {
1108        self.input_dims.len()
1109    }
1110
1111    /// Whether the whole-plan entry would replay in parallel under the
1112    /// active execution policy. Erased callers use it to keep the
1113    /// allocation-free per-segment loop for serial execution.
1114    pub(crate) fn prefers_whole_plan(&self) -> bool {
1115        #[cfg(feature = "parallel")]
1116        {
1117            let total = self.segment_starts[self.segment_starts.len() - 1];
1118            crate::threading::parallel_threads_for_len(total) > 1
1119        }
1120        #[cfg(not(feature = "parallel"))]
1121        {
1122            false
1123        }
1124    }
1125
1126    pub(crate) fn segment_offset(&self, position: usize, dest_offset: isize) -> Result<isize> {
1127        dest_offset
1128            .checked_add(self.dest_offset_deltas[position])
1129            .ok_or(StridedError::OffsetOverflow)
1130    }
1131
1132    pub(crate) fn execute_segment<T>(
1133        &self,
1134        position: usize,
1135        dest: &mut RawStridedMut<'_, T>,
1136        input: &RawStridedRef<'_, T>,
1137    ) -> Result<()>
1138    where
1139        T: Copy + MaybeSendSync,
1140    {
1141        let segment_offset = self.segment_offset(position, dest.offset())?;
1142        let dest_data = dest.data_mut();
1143        let mut segment = unsafe {
1144            RawStridedMut::new_unchecked(
1145                dest_data,
1146                &self.input_dims[position],
1147                &self.dest_strides,
1148                segment_offset,
1149            )
1150        };
1151        self.copy_plans[position].execute(&mut segment, input)
1152    }
1153
1154    pub(crate) fn execute_segment_uninit<T>(
1155        &self,
1156        position: usize,
1157        dest: &mut RawStridedMut<'_, MaybeUninit<T>>,
1158        input: &RawStridedRef<'_, T>,
1159    ) -> Result<()>
1160    where
1161        T: Copy + MaybeSendSync,
1162    {
1163        let segment_offset = self.segment_offset(position, dest.offset())?;
1164        let dest_data = dest.data_mut();
1165        let mut segment = unsafe {
1166            RawStridedMut::new_unchecked(
1167                dest_data,
1168                &self.input_dims[position],
1169                &self.dest_strides,
1170                segment_offset,
1171            )
1172        };
1173        self.copy_plans[position].execute_uninit(&mut segment, input)
1174    }
1175}
1176
1177impl ReversePlan {
1178    /// Compile a reverse plan for one operand layout, destination layout, and axis set.
1179    pub fn compile(
1180        operand_dims: &[usize],
1181        operand_strides: &[isize],
1182        dest_strides: &[isize],
1183        axes: &[usize],
1184    ) -> Result<Self> {
1185        let rank = operand_dims.len();
1186        if operand_strides.len() != rank || dest_strides.len() != rank {
1187            return Err(StridedError::StrideLengthMismatch);
1188        }
1189        checked_total_len(operand_dims)?;
1190
1191        let mut reverse_axis: AxisVec<bool> = (0..rank).map(|_| false).collect();
1192        for &axis in axes {
1193            if axis >= rank {
1194                return Err(StridedError::InvalidAxis { axis, rank });
1195            }
1196            reverse_axis[axis] = true;
1197        }
1198
1199        let mut source_strides: AxisVec<isize> = AxisVec::with_capacity(rank);
1200        let mut source_offset_delta = 0isize;
1201        for axis in 0..rank {
1202            if reverse_axis[axis] {
1203                source_strides.push(
1204                    operand_strides[axis]
1205                        .checked_neg()
1206                        .ok_or(StridedError::OffsetOverflow)?,
1207                );
1208                if operand_dims[axis] > 0 {
1209                    source_offset_delta = checked_offset_add(
1210                        source_offset_delta,
1211                        operand_strides[axis],
1212                        operand_dims[axis] - 1,
1213                    )?;
1214                }
1215            } else {
1216                source_strides.push(operand_strides[axis]);
1217            }
1218        }
1219        let copy_plan = CopyPlan::compile(operand_dims, dest_strides, &source_strides)?;
1220
1221        Ok(Self {
1222            operand_dims: operand_dims.into(),
1223            operand_strides: operand_strides.into(),
1224            dest_strides: dest_strides.into(),
1225            source_strides,
1226            source_offset_delta,
1227            copy_plan,
1228        })
1229    }
1230
1231    /// Execute the prepared reverse traversal.
1232    pub fn execute<T>(
1233        &self,
1234        dest: &mut RawStridedMut<'_, T>,
1235        operand: &RawStridedRef<'_, T>,
1236    ) -> Result<()>
1237    where
1238        T: Copy + MaybeSendSync,
1239    {
1240        self.check_call(dest, operand)?;
1241        let source_offset = operand
1242            .offset()
1243            .checked_add(self.source_offset_delta)
1244            .ok_or(StridedError::OffsetOverflow)?;
1245        let source = unsafe {
1246            RawStridedRef::new_unchecked(
1247                operand.data(),
1248                &self.operand_dims,
1249                &self.source_strides,
1250                source_offset,
1251            )
1252        };
1253        self.copy_plan.execute(dest, &source)
1254    }
1255
1256    /// Execute reverse into storage whose reachable destination elements are uninitialized.
1257    pub fn execute_uninit<T>(
1258        &self,
1259        dest: &mut RawStridedMut<'_, MaybeUninit<T>>,
1260        operand: &RawStridedRef<'_, T>,
1261    ) -> Result<()>
1262    where
1263        T: Copy + MaybeSendSync,
1264    {
1265        self.check_call(dest, operand)?;
1266        let source_offset = operand
1267            .offset()
1268            .checked_add(self.source_offset_delta)
1269            .ok_or(StridedError::OffsetOverflow)?;
1270        let source = unsafe {
1271            RawStridedRef::new_unchecked(
1272                operand.data(),
1273                &self.operand_dims,
1274                &self.source_strides,
1275                source_offset,
1276            )
1277        };
1278        self.copy_plan.execute_uninit(dest, &source)
1279    }
1280
1281    fn check_call<D, T>(
1282        &self,
1283        dest: &RawStridedMut<'_, D>,
1284        operand: &RawStridedRef<'_, T>,
1285    ) -> Result<()> {
1286        if operand.dims() != &self.operand_dims[..]
1287            || operand.strides() != &self.operand_strides[..]
1288            || dest.dims() != &self.operand_dims[..]
1289            || dest.strides() != &self.dest_strides[..]
1290        {
1291            return Err(StridedError::PlanLayoutMismatch);
1292        }
1293        Ok(())
1294    }
1295}
1296
1297fn checked_total_len(dims: &[usize]) -> Result<usize> {
1298    if dims.is_empty() {
1299        return Ok(1);
1300    }
1301    dims.iter()
1302        .try_fold(1usize, |acc, &dim| acc.checked_mul(dim))
1303        .ok_or(StridedError::OffsetOverflow)
1304}
1305
1306fn checked_stride_mul(stride: isize, factor: usize) -> Result<isize> {
1307    let factor = isize::try_from(factor).map_err(|_| StridedError::OffsetOverflow)?;
1308    stride
1309        .checked_mul(factor)
1310        .ok_or(StridedError::OffsetOverflow)
1311}
1312
1313fn checked_pad_output_dim(
1314    input_extent: usize,
1315    edge_low: i64,
1316    edge_high: i64,
1317    interior_step: i64,
1318    axis: usize,
1319    rank: usize,
1320) -> Result<usize> {
1321    let base = if input_extent == 0 {
1322        0i128
1323    } else {
1324        (input_extent as i128 - 1)
1325            .checked_mul(i128::from(interior_step))
1326            .and_then(|value| value.checked_add(1))
1327            .ok_or(StridedError::OffsetOverflow)?
1328    };
1329    let dim = i128::from(edge_low)
1330        .checked_add(i128::from(edge_high))
1331        .and_then(|value| value.checked_add(base))
1332        .ok_or(StridedError::OffsetOverflow)?;
1333    usize::try_from(dim).map_err(|_| StridedError::InvalidAxis { axis, rank })
1334}
1335
1336fn compile_pad_fill_cursor(dims: &[usize], strides: &[isize]) -> Result<PadFillCursor> {
1337    let mut steps = AxisVec::with_capacity(dims.len());
1338    let mut resets = AxisVec::with_capacity(dims.len());
1339    for (&dim, &stride) in dims.iter().zip(strides) {
1340        steps.push(stride);
1341        resets.push(checked_cursor_reset(stride, dim)?);
1342    }
1343    check_offset_span(dims, &steps)?;
1344    Ok(PadFillCursor { steps, resets })
1345}
1346
1347fn compile_pad_copy_cursor(
1348    operand_dims: &[usize],
1349    operand_strides: &[isize],
1350    dest_dims: &[usize],
1351    dest_strides: &[isize],
1352    edge_padding_low: &[i64],
1353    interior_step: &[i64],
1354) -> Result<PadCopyCursor> {
1355    let mut shape = AxisVec::with_capacity(operand_dims.len());
1356    let mut source_steps = AxisVec::with_capacity(operand_dims.len());
1357    let mut source_resets = AxisVec::with_capacity(operand_dims.len());
1358    let mut dest_steps = AxisVec::with_capacity(operand_dims.len());
1359    let mut dest_resets = AxisVec::with_capacity(operand_dims.len());
1360    let mut source_base_delta = 0isize;
1361    let mut dest_base_delta = 0isize;
1362    let mut copy_empty = false;
1363
1364    for axis in 0..operand_dims.len() {
1365        let (start, end) = checked_pad_valid_interval(
1366            operand_dims[axis],
1367            dest_dims[axis],
1368            edge_padding_low[axis],
1369            interior_step[axis],
1370        )?;
1371        let extent = end - start;
1372        shape.push(extent);
1373        if !copy_empty && extent != 0 {
1374            source_base_delta =
1375                checked_offset_add(source_base_delta, operand_strides[axis], start)?;
1376            let output_start = i128::from(edge_padding_low[axis])
1377                .checked_add(
1378                    i128::try_from(start)
1379                        .map_err(|_| StridedError::OffsetOverflow)?
1380                        .checked_mul(i128::from(interior_step[axis]))
1381                        .ok_or(StridedError::OffsetOverflow)?,
1382                )
1383                .ok_or(StridedError::OffsetOverflow)?;
1384            let output_start =
1385                usize::try_from(output_start).map_err(|_| StridedError::OffsetOverflow)?;
1386            dest_base_delta =
1387                checked_offset_add(dest_base_delta, dest_strides[axis], output_start)?;
1388        }
1389        copy_empty |= extent == 0;
1390
1391        let source_step = operand_strides[axis];
1392        let dest_step = checked_stride_mul_i64(dest_strides[axis], interior_step[axis])?;
1393        source_steps.push(source_step);
1394        source_resets.push(checked_cursor_reset(source_step, extent)?);
1395        dest_steps.push(dest_step);
1396        dest_resets.push(checked_cursor_reset(dest_step, extent)?);
1397    }
1398
1399    check_offset_span(&shape, &source_steps)?;
1400    check_offset_span(&shape, &dest_steps)?;
1401    let total = if shape.iter().any(|&extent| extent == 0) {
1402        0
1403    } else {
1404        checked_total_len(&shape)?
1405    };
1406
1407    Ok(PadCopyCursor {
1408        shape,
1409        source_base_delta,
1410        dest_base_delta,
1411        source_steps,
1412        source_resets,
1413        dest_steps,
1414        dest_resets,
1415        total,
1416    })
1417}
1418
1419fn checked_pad_valid_interval(
1420    input_extent: usize,
1421    dest_extent: usize,
1422    edge_low: i64,
1423    step: i64,
1424) -> Result<(usize, usize)> {
1425    if input_extent == 0 || dest_extent == 0 {
1426        return Ok((0, 0));
1427    }
1428    let step = i128::from(step);
1429    let lower = ceil_div_positive(-i128::from(edge_low), step)?;
1430    let dest_last = i128::try_from(dest_extent)
1431        .map_err(|_| StridedError::OffsetOverflow)?
1432        .checked_sub(1)
1433        .ok_or(StridedError::OffsetOverflow)?;
1434    let upper = floor_div_positive(
1435        dest_last
1436            .checked_sub(i128::from(edge_low))
1437            .ok_or(StridedError::OffsetOverflow)?,
1438        step,
1439    )?
1440    .checked_add(1)
1441    .ok_or(StridedError::OffsetOverflow)?;
1442    let input_extent = i128::try_from(input_extent).map_err(|_| StridedError::OffsetOverflow)?;
1443    let lower = lower.clamp(0, input_extent);
1444    let upper = upper.clamp(0, input_extent);
1445    if lower >= upper {
1446        return Ok((0, 0));
1447    }
1448    Ok((
1449        usize::try_from(lower).map_err(|_| StridedError::OffsetOverflow)?,
1450        usize::try_from(upper).map_err(|_| StridedError::OffsetOverflow)?,
1451    ))
1452}
1453
1454fn ceil_div_positive(numerator: i128, denominator: i128) -> Result<i128> {
1455    if denominator <= 0 {
1456        return Err(StridedError::OffsetOverflow);
1457    }
1458    let quotient = numerator.div_euclid(denominator);
1459    let remainder = numerator.rem_euclid(denominator);
1460    quotient
1461        .checked_add(i128::from(remainder != 0))
1462        .ok_or(StridedError::OffsetOverflow)
1463}
1464
1465fn floor_div_positive(numerator: i128, denominator: i128) -> Result<i128> {
1466    if denominator <= 0 {
1467        return Err(StridedError::OffsetOverflow);
1468    }
1469    Ok(numerator.div_euclid(denominator))
1470}
1471
1472fn checked_stride_mul_i64(stride: isize, factor: i64) -> Result<isize> {
1473    let factor = isize::try_from(factor).map_err(|_| StridedError::OffsetOverflow)?;
1474    stride
1475        .checked_mul(factor)
1476        .ok_or(StridedError::OffsetOverflow)
1477}
1478
1479fn checked_cursor_reset(step: isize, extent: usize) -> Result<isize> {
1480    if extent == 0 {
1481        return Ok(0);
1482    }
1483    let last = isize::try_from(extent - 1).map_err(|_| StridedError::OffsetOverflow)?;
1484    step.checked_mul(last)
1485        .and_then(isize::checked_neg)
1486        .ok_or(StridedError::OffsetOverflow)
1487}
1488
1489fn check_offset_span(shape: &[usize], strides: &[isize]) -> Result<()> {
1490    let mut min = 0isize;
1491    let mut max = 0isize;
1492    for (&extent, &stride) in shape.iter().zip(strides) {
1493        if extent <= 1 {
1494            continue;
1495        }
1496        let last = isize::try_from(extent - 1).map_err(|_| StridedError::OffsetOverflow)?;
1497        let delta = stride
1498            .checked_mul(last)
1499            .ok_or(StridedError::OffsetOverflow)?;
1500        if delta < 0 {
1501            min = min.checked_add(delta).ok_or(StridedError::OffsetOverflow)?;
1502        } else {
1503            max = max.checked_add(delta).ok_or(StridedError::OffsetOverflow)?;
1504        }
1505    }
1506    let _ = (min, max);
1507    Ok(())
1508}
1509
1510fn is_dense_col_major(dims: &[usize], strides: &[isize]) -> bool {
1511    let mut expected = 1isize;
1512    for (&dim, &stride) in dims.iter().zip(strides.iter()) {
1513        if stride != expected {
1514            return false;
1515        }
1516        let Ok(dim) = isize::try_from(dim) else {
1517            return false;
1518        };
1519        let Some(next) = expected.checked_mul(dim) else {
1520            return false;
1521        };
1522        expected = next;
1523    }
1524    true
1525}
1526
1527fn compile_contiguous_pad_axis0_run(
1528    operand_dims: &[usize],
1529    operand_strides: &[isize],
1530    dest_dims: &[usize],
1531    dest_strides: &[isize],
1532    edge_padding_low: &[i64],
1533    interior_step: &[i64],
1534) -> Option<ContiguousPadAxis0Run> {
1535    if operand_dims.is_empty()
1536        || operand_strides[0] != 1
1537        || dest_strides[0] != 1
1538        || interior_step[0] != 1
1539    {
1540        return None;
1541    }
1542
1543    let operand_extent = operand_dims[0] as i128;
1544    let dest_extent = dest_dims[0] as i128;
1545    let edge_low = i128::from(edge_padding_low[0]);
1546    let operand_start = (-edge_low).clamp(0, operand_extent);
1547    let dest_start = edge_low.clamp(0, dest_extent);
1548    let len = (operand_extent - operand_start).min(dest_extent - dest_start);
1549    Some(ContiguousPadAxis0Run {
1550        operand_start: usize::try_from(operand_start).ok()?,
1551        dest_start: usize::try_from(dest_start).ok()?,
1552        len: usize::try_from(len).ok()?,
1553    })
1554}
1555
1556fn checked_offset_add(base: isize, stride: isize, coord: usize) -> Result<isize> {
1557    let coord = isize::try_from(coord).map_err(|_| StridedError::OffsetOverflow)?;
1558    let scaled = stride
1559        .checked_mul(coord)
1560        .ok_or(StridedError::OffsetOverflow)?;
1561    base.checked_add(scaled).ok_or(StridedError::OffsetOverflow)
1562}
1563
1564fn advance_col_major_index(index: &mut [usize], shape: &[usize]) {
1565    for axis in 0..index.len() {
1566        index[axis] += 1;
1567        if index[axis] < shape[axis] {
1568            return;
1569        }
1570        index[axis] = 0;
1571    }
1572}
1573
1574#[cfg(feature = "parallel")]
1575fn fill_col_major_index(mut linear: usize, shape: &[usize], out: &mut [usize]) {
1576    for (axis, coord) in out.iter_mut().enumerate() {
1577        let dim = shape[axis];
1578        *coord = linear % dim;
1579        linear /= dim;
1580    }
1581}
1582
1583struct PadFillState {
1584    coords: CoordScratch,
1585    offset: isize,
1586}
1587
1588impl PadFillState {
1589    fn new(base: isize, shape: &[usize], _cursor: &PadFillCursor) -> Self {
1590        Self {
1591            coords: CoordScratch::new(shape.len()),
1592            offset: base,
1593        }
1594    }
1595
1596    #[cfg(feature = "parallel")]
1597    fn decode(linear: usize, base: isize, shape: &[usize], cursor: &PadFillCursor) -> Result<Self> {
1598        let mut state = Self::new(base, shape, cursor);
1599        fill_col_major_index(linear, shape, state.coords.as_mut_slice());
1600        for (&coord, &step) in state.coords.as_mut_slice().iter().zip(&cursor.steps) {
1601            state.offset = checked_offset_add(state.offset, step, coord)?;
1602        }
1603        Ok(state)
1604    }
1605
1606    #[inline]
1607    fn advance(&mut self, shape: &[usize], cursor: &PadFillCursor) {
1608        // INVARIANT: precomputed checked offset arithmetic, the validated
1609        // descriptor, and exact layout equality prove each incremental offset
1610        // remains reachable.
1611        for axis in 0..shape.len() {
1612            let next = self.coords.as_mut_slice()[axis] + 1;
1613            if next < shape[axis] {
1614                self.coords.as_mut_slice()[axis] = next;
1615                self.offset += cursor.steps[axis];
1616                return;
1617            }
1618            self.coords.as_mut_slice()[axis] = 0;
1619            self.offset += cursor.resets[axis];
1620        }
1621    }
1622}
1623
1624struct PadCopyState {
1625    coords: CoordScratch,
1626    source_offset: isize,
1627    dest_offset: isize,
1628}
1629
1630impl PadCopyState {
1631    fn new(source_base: isize, dest_base: isize, cursor: &PadCopyCursor) -> Self {
1632        Self {
1633            coords: CoordScratch::new(cursor.shape.len()),
1634            source_offset: source_base,
1635            dest_offset: dest_base,
1636        }
1637    }
1638
1639    #[cfg(feature = "parallel")]
1640    fn decode(
1641        linear: usize,
1642        source_base: isize,
1643        dest_base: isize,
1644        cursor: &PadCopyCursor,
1645    ) -> Result<Self> {
1646        let mut state = Self::new(source_base, dest_base, cursor);
1647        fill_col_major_index(linear, &cursor.shape, state.coords.as_mut_slice());
1648        for axis in 0..cursor.shape.len() {
1649            let coord = state.coords.as_mut_slice()[axis];
1650            state.source_offset =
1651                checked_offset_add(state.source_offset, cursor.source_steps[axis], coord)?;
1652            state.dest_offset =
1653                checked_offset_add(state.dest_offset, cursor.dest_steps[axis], coord)?;
1654        }
1655        Ok(state)
1656    }
1657
1658    #[inline]
1659    fn advance(&mut self, cursor: &PadCopyCursor) {
1660        // INVARIANT: compile-time checked steps/resets/offset arithmetic,
1661        // validated descriptors, and exact plan-layout equality prove every update below
1662        // stays in the reachable source/destination domains.
1663        for axis in 0..cursor.shape.len() {
1664            let next = self.coords.as_mut_slice()[axis] + 1;
1665            if next < cursor.shape[axis] {
1666                self.coords.as_mut_slice()[axis] = next;
1667                self.source_offset += cursor.source_steps[axis];
1668                self.dest_offset += cursor.dest_steps[axis];
1669                return;
1670            }
1671            self.coords.as_mut_slice()[axis] = 0;
1672            self.source_offset += cursor.source_resets[axis];
1673            self.dest_offset += cursor.dest_resets[axis];
1674        }
1675    }
1676}
1677
1678struct CoordScratch {
1679    inline: [usize; crate::RAW_FUSED_RANK_LIMIT],
1680    heap: Option<Vec<usize>>,
1681    len: usize,
1682}
1683
1684impl CoordScratch {
1685    fn new(len: usize) -> Self {
1686        if len <= crate::RAW_FUSED_RANK_LIMIT {
1687            Self {
1688                inline: [0; crate::RAW_FUSED_RANK_LIMIT],
1689                heap: None,
1690                len,
1691            }
1692        } else {
1693            Self {
1694                inline: [0; crate::RAW_FUSED_RANK_LIMIT],
1695                heap: Some(vec![0; len]),
1696                len,
1697            }
1698        }
1699    }
1700
1701    fn as_mut_slice(&mut self) -> &mut [usize] {
1702        match &mut self.heap {
1703            Some(heap) => heap,
1704            None => &mut self.inline[..self.len],
1705        }
1706    }
1707}
1708
1709#[cfg(test)]
1710#[path = "static_indexing_plan/tests/tests.rs"]
1711mod tests;
1712
1713#[cfg(all(test, feature = "parallel"))]
1714#[path = "static_indexing_plan/tests/parallel_tests.rs"]
1715mod parallel_tests;