Skip to main content

strided_basic/
gather_plan.rs

1//! Prepared indexed plans over raw strided value and index layouts.
2//!
3//! This module owns the generic gather, dynamic-slice/update, and scatter
4//! traversals used by the erased replay layer. It models the XLA/tenferro
5//! indexed shape vocabulary, but keeps tensor allocation, dtype promotion, and
6//! frontend error policy outside `strided-kernel`.
7
8use core::{mem::MaybeUninit, ops::Add};
9
10use crate::copy_plan::{CopyPlan, OverwriteWriter, ReadModifyWrite};
11use crate::{
12    MaybeSendSync, RawStridedMut, RawStridedRef, Result, StridedError, RAW_FUSED_RANK_LIMIT,
13};
14
15#[cfg(feature = "parallel")]
16type AxisVec<T> = smallvec::SmallVec<[T; RAW_FUSED_RANK_LIMIT]>;
17#[cfg(not(feature = "parallel"))]
18type AxisVec<T> = Vec<T>;
19
20/// Gather configuration shared by generic and erased replay.
21///
22/// The fields follow the usual gather vocabulary:
23///
24/// - `start_index_map[component]` names the operand axis controlled by a
25///   component in the index vector;
26/// - `collapsed_slice_dims` names operand axes whose slice size is one and
27///   which do not appear as output window axes;
28/// - `offset_dims` names output axes that represent window offsets;
29/// - output axes not in `offset_dims` are batch axes from `start_indices`;
30/// - `index_vector_dim == start_indices_rank` represents scalar index vectors.
31#[derive(Clone, Debug, Eq, PartialEq)]
32pub struct GatherSpec {
33    pub offset_dims: Vec<usize>,
34    pub collapsed_slice_dims: Vec<usize>,
35    pub start_index_map: Vec<usize>,
36    pub index_vector_dim: usize,
37    pub slice_sizes: Vec<usize>,
38}
39
40/// Index scalar types accepted by [`GatherPlan`].
41pub trait GatherIndex: Copy + MaybeSendSync {
42    fn to_i64(self) -> i64;
43}
44
45impl GatherIndex for i32 {
46    #[inline]
47    fn to_i64(self) -> i64 {
48        i64::from(self)
49    }
50}
51
52impl GatherIndex for i64 {
53    #[inline]
54    fn to_i64(self) -> i64 {
55        self
56    }
57}
58
59/// A compiled gather traversal for one value layout, index layout, and output layout.
60#[derive(Clone, Debug)]
61pub struct GatherPlan {
62    operand_dims: AxisVec<usize>,
63    operand_strides: AxisVec<isize>,
64    index_dims: AxisVec<usize>,
65    index_strides: AxisVec<isize>,
66    dest_dims: AxisVec<usize>,
67    dest_strides: AxisVec<isize>,
68    spec: GatherSpec,
69    replay: GatherReplay,
70    total: usize,
71}
72
73#[derive(Clone, Copy, Debug)]
74struct GatherReplayAxis {
75    dest_step: isize,
76    dest_reset: isize,
77    window_step: isize,
78    window_reset: isize,
79    index_batch_step: isize,
80    index_batch_reset: isize,
81}
82
83#[derive(Clone, Debug)]
84struct GatherReplay {
85    axes: AxisVec<GatherReplayAxis>,
86    index_component_offsets: AxisVec<isize>,
87    index_operand_strides: AxisVec<isize>,
88}
89
90struct GatherReplayState {
91    coords: AxisVec<usize>,
92    dest_offset: isize,
93    window_offset: isize,
94    index_batch_offset: isize,
95}
96
97#[derive(Clone, Copy, Debug)]
98struct WindowReplayAxis {
99    source_step: isize,
100    source_reset: isize,
101    dest_step: isize,
102    dest_reset: isize,
103}
104
105#[derive(Clone, Debug)]
106struct WindowReplay {
107    shape: AxisVec<usize>,
108    axes: AxisVec<WindowReplayAxis>,
109}
110
111struct WindowReplayState {
112    coords: CoordScratch,
113    source_offset: isize,
114    dest_offset: isize,
115}
116
117#[derive(Clone, Debug)]
118struct ScatterReplay {
119    batch: WindowReplay,
120    window: WindowReplay,
121    index_component_offsets: AxisVec<isize>,
122}
123
124impl ScatterReplay {
125    fn compile(
126        batch_shape: &[usize],
127        index_dims: &[usize],
128        index_strides: &[isize],
129        index_vector_dim: usize,
130        index_component_count: usize,
131        update_dims: &[usize],
132        update_strides: &[isize],
133        update_window_dims: &[usize],
134        is_update_window_dim: &[bool],
135        window_shape_updates: &[usize],
136        window_dims: &[usize],
137        dest_strides: &[isize],
138    ) -> Result<Self> {
139        let (
140            batch_source_strides,
141            batch_dest_strides,
142            window_source_strides,
143            window_dest_strides,
144            vector_stride,
145        ) = scatter_replay_strides(
146            index_dims,
147            index_strides,
148            index_vector_dim,
149            update_dims,
150            update_strides,
151            update_window_dims,
152            is_update_window_dim,
153            window_dims,
154            dest_strides,
155        );
156        let batch = WindowReplay::compile(batch_shape, &batch_source_strides, &batch_dest_strides)?;
157        let window = WindowReplay::compile(
158            window_shape_updates,
159            &window_source_strides,
160            &window_dest_strides,
161        )?;
162        let mut index_component_offsets = AxisVec::with_capacity(index_component_count);
163        for component in 0..index_component_count {
164            index_component_offsets.push(checked_offset_add(0, vector_stride, component)?);
165        }
166        Ok(Self {
167            batch,
168            window,
169            index_component_offsets,
170        })
171    }
172
173    #[cfg(feature = "parallel")]
174    fn validate(
175        batch_shape: &[usize],
176        index_dims: &[usize],
177        index_strides: &[isize],
178        index_vector_dim: usize,
179        index_component_count: usize,
180        update_dims: &[usize],
181        update_strides: &[isize],
182        update_window_dims: &[usize],
183        is_update_window_dim: &[bool],
184        window_shape_updates: &[usize],
185        window_dims: &[usize],
186        dest_strides: &[isize],
187    ) -> Result<()> {
188        let (
189            batch_source_strides,
190            batch_dest_strides,
191            window_source_strides,
192            window_dest_strides,
193            vector_stride,
194        ) = scatter_replay_strides(
195            index_dims,
196            index_strides,
197            index_vector_dim,
198            update_dims,
199            update_strides,
200            update_window_dims,
201            is_update_window_dim,
202            window_dims,
203            dest_strides,
204        );
205        WindowReplay::compile(batch_shape, &batch_source_strides, &batch_dest_strides)?;
206        WindowReplay::compile(
207            window_shape_updates,
208            &window_source_strides,
209            &window_dest_strides,
210        )?;
211        for component in 0..index_component_count {
212            checked_offset_add(0, vector_stride, component)?;
213        }
214        Ok(())
215    }
216}
217
218fn scatter_replay_strides(
219    index_dims: &[usize],
220    index_strides: &[isize],
221    index_vector_dim: usize,
222    update_dims: &[usize],
223    update_strides: &[isize],
224    update_window_dims: &[usize],
225    is_update_window_dim: &[bool],
226    window_dims: &[usize],
227    dest_strides: &[isize],
228) -> (
229    AxisVec<isize>,
230    AxisVec<isize>,
231    AxisVec<isize>,
232    AxisVec<isize>,
233    isize,
234) {
235    let batch_source_strides = index_dims
236        .iter()
237        .zip(index_strides.iter())
238        .enumerate()
239        .filter_map(|(axis, (_, &stride))| (axis != index_vector_dim).then_some(stride))
240        .collect();
241    let batch_dest_strides = update_dims
242        .iter()
243        .zip(update_strides.iter())
244        .enumerate()
245        .filter_map(|(axis, (_, &stride))| (!is_update_window_dim[axis]).then_some(stride))
246        .collect();
247    let window_source_strides = update_window_dims
248        .iter()
249        .map(|&axis| update_strides[axis])
250        .collect();
251    let window_dest_strides = window_dims.iter().map(|&axis| dest_strides[axis]).collect();
252    let vector_stride = if index_vector_dim < index_dims.len() {
253        index_strides[index_vector_dim]
254    } else {
255        0
256    };
257    (
258        batch_source_strides,
259        batch_dest_strides,
260        window_source_strides,
261        window_dest_strides,
262        vector_stride,
263    )
264}
265
266impl WindowReplay {
267    fn compile(shape: &[usize], source_strides: &[isize], dest_strides: &[isize]) -> Result<Self> {
268        validate_layout_span(shape, source_strides)?;
269        validate_layout_span(shape, dest_strides)?;
270        let mut fused_shape: AxisVec<usize> = AxisVec::with_capacity(shape.len());
271        let mut axes: AxisVec<WindowReplayAxis> = AxisVec::with_capacity(shape.len());
272        for (axis, &dim) in shape.iter().enumerate() {
273            let source_step = source_strides[axis];
274            let dest_step = dest_strides[axis];
275            if let Some(previous_axis) = fused_shape.len().checked_sub(1) {
276                let previous_extent = fused_shape[previous_axis];
277                let previous_extent =
278                    isize::try_from(previous_extent).map_err(|_| StridedError::OffsetOverflow)?;
279                let expected_source = axes[previous_axis]
280                    .source_step
281                    .checked_mul(previous_extent)
282                    .ok_or(StridedError::OffsetOverflow)?;
283                let expected_dest = axes[previous_axis]
284                    .dest_step
285                    .checked_mul(previous_extent)
286                    .ok_or(StridedError::OffsetOverflow)?;
287                if source_step == expected_source && dest_step == expected_dest {
288                    let fused_extent = fused_shape[previous_axis]
289                        .checked_mul(dim)
290                        .ok_or(StridedError::OffsetOverflow)?;
291                    fused_shape[previous_axis] = fused_extent;
292                    axes[previous_axis].source_reset =
293                        checked_replay_reset(fused_extent, axes[previous_axis].source_step)?;
294                    axes[previous_axis].dest_reset =
295                        checked_replay_reset(fused_extent, axes[previous_axis].dest_step)?;
296                    continue;
297                }
298            }
299            fused_shape.push(dim);
300            axes.push(WindowReplayAxis {
301                source_step,
302                source_reset: checked_replay_reset(dim, source_step)?,
303                dest_step,
304                dest_reset: checked_replay_reset(dim, dest_step)?,
305            });
306        }
307        Ok(Self {
308            shape: fused_shape,
309            axes,
310        })
311    }
312
313    fn decode(
314        &self,
315        mut linear: usize,
316        source_base: isize,
317        dest_base: isize,
318    ) -> Result<WindowReplayState> {
319        let mut coords = CoordScratch::new(self.shape.len());
320        let mut source_offset = source_base;
321        let mut dest_offset = dest_base;
322        for (axis, (&dim, coord)) in self
323            .shape
324            .iter()
325            .zip(coords.as_mut_slice().iter_mut())
326            .enumerate()
327        {
328            *coord = linear % dim;
329            linear /= dim;
330            let replay_axis = self.axes[axis];
331            source_offset = checked_offset_add(source_offset, replay_axis.source_step, *coord)?;
332            dest_offset = checked_offset_add(dest_offset, replay_axis.dest_step, *coord)?;
333        }
334        Ok(WindowReplayState {
335            coords,
336            source_offset,
337            dest_offset,
338        })
339    }
340
341    #[inline]
342    fn advance(&self, state: &mut WindowReplayState) {
343        for ((coord, &dim), replay_axis) in state
344            .coords
345            .as_mut_slice()
346            .iter_mut()
347            .zip(self.shape.iter())
348            .zip(self.axes.iter())
349        {
350            let next = *coord + 1;
351            if next < dim {
352                *coord = next;
353                state.source_offset += replay_axis.source_step;
354                state.dest_offset += replay_axis.dest_step;
355                return;
356            }
357            *coord = 0;
358            state.source_offset += replay_axis.source_reset;
359            state.dest_offset += replay_axis.dest_reset;
360        }
361    }
362}
363
364/// Scatter configuration shared by generic and erased replay.
365///
366/// `ScatterPlan` implements tenferro's current additive scatter semantics:
367/// every update value is added to the selected output slot, so overlapping
368/// windows accumulate in deterministic column-major replay order.
369#[derive(Clone, Debug, Eq, PartialEq)]
370pub struct ScatterSpec {
371    pub update_window_dims: Vec<usize>,
372    pub inserted_window_dims: Vec<usize>,
373    pub scatter_dims_to_operand_dims: Vec<usize>,
374    pub index_vector_dim: usize,
375}
376
377/// A compiled fixed-window dynamic-slice traversal.
378#[derive(Clone, Debug)]
379pub struct DynamicSlicePlan {
380    operand_dims: AxisVec<usize>,
381    operand_strides: AxisVec<isize>,
382    start_dims: AxisVec<usize>,
383    start_strides: AxisVec<isize>,
384    dest_dims: AxisVec<usize>,
385    dest_strides: AxisVec<isize>,
386    slice_sizes: AxisVec<usize>,
387    total: usize,
388    #[cfg(not(feature = "parallel"))]
389    replay: WindowReplay,
390}
391
392/// A compiled dynamic-update-slice traversal.
393///
394/// Execution first copies `operand` into `dest`, then overwrites the clamped
395/// update window. The plan performs no allocation for ranks at most
396/// [`RAW_FUSED_RANK_LIMIT`].
397#[derive(Clone, Debug)]
398pub struct DynamicUpdateSlicePlan {
399    operand_dims: AxisVec<usize>,
400    operand_strides: AxisVec<isize>,
401    start_dims: AxisVec<usize>,
402    start_strides: AxisVec<isize>,
403    update_dims: AxisVec<usize>,
404    update_strides: AxisVec<isize>,
405    dest_dims: AxisVec<usize>,
406    dest_strides: AxisVec<isize>,
407    total: usize,
408    copy_plan: CopyPlan,
409    #[cfg(not(feature = "parallel"))]
410    replay: WindowReplay,
411}
412
413/// A compiled additive scatter traversal.
414///
415/// Execution first copies `operand` into `dest`, then applies additive updates
416/// in deterministic column-major order. Boolean values are intentionally not
417/// supported because additive scatter has no bool semantics.
418#[derive(Clone, Debug)]
419pub struct ScatterPlan {
420    operand_dims: AxisVec<usize>,
421    operand_strides: AxisVec<isize>,
422    index_dims: AxisVec<usize>,
423    index_strides: AxisVec<isize>,
424    update_dims: AxisVec<usize>,
425    update_strides: AxisVec<isize>,
426    dest_dims: AxisVec<usize>,
427    dest_strides: AxisVec<isize>,
428    spec: ScatterSpec,
429    #[cfg(feature = "parallel")]
430    batch_shape: AxisVec<usize>,
431    #[cfg(feature = "parallel")]
432    window_dims: AxisVec<usize>,
433    window_shape: AxisVec<usize>,
434    #[cfg(feature = "parallel")]
435    window_shape_updates: AxisVec<usize>,
436    #[cfg(feature = "parallel")]
437    is_update_window_dim: AxisVec<bool>,
438    batch_elems: usize,
439    window_elems: usize,
440    copy_plan: CopyPlan,
441    #[cfg(not(feature = "parallel"))]
442    replay: ScatterReplay,
443}
444
445impl GatherPlan {
446    /// Compile a gather plan for fixed operand, index, and destination layouts.
447    pub fn compile(
448        operand_dims: &[usize],
449        operand_strides: &[isize],
450        index_dims: &[usize],
451        index_strides: &[isize],
452        dest_dims: &[usize],
453        dest_strides: &[isize],
454        spec: GatherSpec,
455    ) -> Result<Self> {
456        if operand_dims.len() != operand_strides.len()
457            || index_dims.len() != index_strides.len()
458            || dest_dims.len() != dest_strides.len()
459        {
460            return Err(StridedError::StrideLengthMismatch);
461        }
462        checked_total_len(operand_dims)?;
463        checked_total_len(index_dims)?;
464        let total = checked_total_len(dest_dims)?;
465        validate_layout_span(operand_dims, operand_strides)?;
466        validate_layout_span(index_dims, index_strides)?;
467        validate_layout_span(dest_dims, dest_strides)?;
468        if !crate::layout_check::is_injective_layout(dest_dims, dest_strides) {
469            return Err(StridedError::NonInjectiveOutputLayout);
470        }
471
472        let operand_rank = operand_dims.len();
473        if spec.slice_sizes.len() != operand_rank {
474            return Err(StridedError::RankMismatch(
475                spec.slice_sizes.len(),
476                operand_rank,
477            ));
478        }
479        validate_unique_axes(&spec.collapsed_slice_dims, operand_rank)?;
480        validate_unique_axes(&spec.start_index_map, operand_rank)?;
481        if spec.index_vector_dim > index_dims.len() {
482            return Err(StridedError::InvalidAxis {
483                axis: spec.index_vector_dim,
484                rank: index_dims.len() + 1,
485            });
486        }
487
488        for (axis, (&window, &dim)) in spec.slice_sizes.iter().zip(operand_dims.iter()).enumerate()
489        {
490            if window > dim {
491                return Err(StridedError::InvalidAxis {
492                    axis,
493                    rank: operand_rank,
494                });
495            }
496        }
497        for &axis in &spec.collapsed_slice_dims {
498            if spec.slice_sizes[axis] != 1 {
499                return Err(StridedError::InvalidAxis {
500                    axis,
501                    rank: operand_rank,
502                });
503            }
504        }
505
506        let index_vector_size = if spec.index_vector_dim == index_dims.len() {
507            1
508        } else {
509            index_dims[spec.index_vector_dim]
510        };
511        if index_vector_size != spec.start_index_map.len() {
512            return Err(StridedError::RankMismatch(
513                index_vector_size,
514                spec.start_index_map.len(),
515            ));
516        }
517
518        let window_dims = operand_window_dims(operand_rank, &spec.collapsed_slice_dims);
519        if spec.offset_dims.len() != window_dims.len() {
520            return Err(StridedError::RankMismatch(
521                spec.offset_dims.len(),
522                window_dims.len(),
523            ));
524        }
525
526        let batch_shape = index_batch_shape(index_dims, spec.index_vector_dim);
527        let out_rank = batch_shape.len() + spec.offset_dims.len();
528        validate_unique_axes(&spec.offset_dims, out_rank)?;
529
530        let mut out_axis_to_operand_dim: AxisVec<Option<usize>> =
531            (0..out_rank).map(|_| None).collect();
532        for (offset_axis, &out_axis) in spec.offset_dims.iter().enumerate() {
533            out_axis_to_operand_dim[out_axis] = Some(window_dims[offset_axis]);
534        }
535
536        let mut expected_dest_dims: AxisVec<usize> = AxisVec::with_capacity(out_rank);
537        let mut batch_axis = 0usize;
538        for &operand_dim in &out_axis_to_operand_dim {
539            match operand_dim {
540                Some(axis) => expected_dest_dims.push(spec.slice_sizes[axis]),
541                None => {
542                    expected_dest_dims.push(batch_shape[batch_axis]);
543                    batch_axis += 1;
544                }
545            }
546        }
547        if dest_dims != &expected_dest_dims[..] {
548            return Err(StridedError::ShapeMismatch(
549                dest_dims.to_vec(),
550                expected_dest_dims.to_vec(),
551            ));
552        }
553
554        let batch_index_strides: AxisVec<isize> = index_dims
555            .iter()
556            .zip(index_strides.iter())
557            .enumerate()
558            .filter_map(|(axis, (_, &stride))| (axis != spec.index_vector_dim).then_some(stride))
559            .collect();
560        let mut batch_axis = 0usize;
561        let mut replay_axes = AxisVec::with_capacity(out_rank);
562        for (out_axis, &operand_dim) in out_axis_to_operand_dim.iter().enumerate() {
563            let (window_step, index_batch_step) = match operand_dim {
564                Some(axis) => (operand_strides[axis], 0),
565                None => {
566                    let step = batch_index_strides[batch_axis];
567                    batch_axis += 1;
568                    (0, step)
569                }
570            };
571            replay_axes.push(GatherReplayAxis {
572                dest_step: dest_strides[out_axis],
573                dest_reset: checked_replay_reset(dest_dims[out_axis], dest_strides[out_axis])?,
574                window_step,
575                window_reset: checked_replay_reset(dest_dims[out_axis], window_step)?,
576                index_batch_step,
577                index_batch_reset: checked_replay_reset(dest_dims[out_axis], index_batch_step)?,
578            });
579        }
580
581        let vector_stride = if spec.index_vector_dim < index_dims.len() {
582            index_strides[spec.index_vector_dim]
583        } else {
584            0
585        };
586        let mut index_component_offsets = AxisVec::with_capacity(spec.start_index_map.len());
587        for component in 0..spec.start_index_map.len() {
588            index_component_offsets.push(checked_offset_add(0, vector_stride, component)?);
589        }
590
591        Ok(Self {
592            operand_dims: operand_dims.into(),
593            operand_strides: operand_strides.into(),
594            index_dims: index_dims.into(),
595            index_strides: index_strides.into(),
596            dest_dims: dest_dims.into(),
597            dest_strides: dest_strides.into(),
598            spec: spec.clone(),
599            replay: GatherReplay {
600                axes: replay_axes,
601                index_component_offsets,
602                index_operand_strides: spec
603                    .start_index_map
604                    .iter()
605                    .map(|&axis| operand_strides[axis])
606                    .collect(),
607            },
608            total,
609        })
610    }
611
612    #[inline]
613    pub fn spec(&self) -> &GatherSpec {
614        &self.spec
615    }
616
617    #[inline]
618    pub fn dest_dims(&self) -> &[usize] {
619        &self.dest_dims
620    }
621
622    /// Execute the prepared gather traversal.
623    pub fn execute<T, I>(
624        &self,
625        dest: &mut RawStridedMut<'_, T>,
626        operand: &RawStridedRef<'_, T>,
627        start_indices: &RawStridedRef<'_, I>,
628    ) -> Result<()>
629    where
630        T: Copy + MaybeSendSync,
631        I: GatherIndex,
632    {
633        self.execute_with_writer(dest, operand, start_indices)
634    }
635
636    /// Execute the prepared gather into a destination whose reachable slots
637    /// may be uninitialized. Every logical destination slot is written.
638    pub(crate) fn execute_uninit<T, I>(
639        &self,
640        dest: &mut RawStridedMut<'_, MaybeUninit<T>>,
641        operand: &RawStridedRef<'_, T>,
642        start_indices: &RawStridedRef<'_, I>,
643    ) -> Result<()>
644    where
645        T: Copy + MaybeSendSync,
646        I: GatherIndex,
647    {
648        self.execute_with_writer(dest, operand, start_indices)
649    }
650
651    fn execute_with_writer<T, I, W>(
652        &self,
653        dest: &mut W,
654        operand: &RawStridedRef<'_, T>,
655        start_indices: &RawStridedRef<'_, I>,
656    ) -> Result<()>
657    where
658        T: Copy + MaybeSendSync,
659        I: GatherIndex,
660        W: OverwriteWriter<T>,
661    {
662        self.check_call(dest, operand, start_indices)?;
663        if self.total == 0 {
664            return Ok(());
665        }
666        if self.uses_rank_one_scalar_take_path() {
667            #[cfg(feature = "parallel")]
668            {
669                let nthreads = crate::threading::parallel_threads_for_len(self.total);
670                if nthreads > 1 {
671                    return self.execute_rank_one_scalar_take_parallel(
672                        dest,
673                        operand,
674                        start_indices,
675                        nthreads,
676                    );
677                }
678            }
679            return self.execute_rank_one_scalar_take(dest, operand, start_indices);
680        }
681        #[cfg(feature = "parallel")]
682        {
683            let nthreads = crate::threading::parallel_threads_for_len(self.total);
684            if nthreads > 1 {
685                return self.execute_parallel(dest, operand, start_indices, nthreads);
686            }
687        }
688
689        let mut state =
690            self.decode_replay_state(0, dest.offset(), operand.offset(), start_indices.offset())?;
691        let operand_data = operand.data();
692        let index_data = start_indices.data();
693
694        // INVARIANT: the checked decode starts at a valid logical output. Each
695        // replay axis has checked step/reset deltas, so this state remains the
696        // corresponding destination, window, and batch-index offsets.
697        for _ in 0..self.total {
698            // INVARIANT: window_offset plus every clamped mapped contribution
699            // is an operand offset for coordinates inside operand_dims; the
700            // validated operand span and RawStridedRef reachability therefore
701            // keep the fresh source offset inside the allocation.
702            let mut source_offset = state.window_offset;
703            for ((&index_component_offset, &operand_stride), &operand_dim) in self
704                .replay
705                .index_component_offsets
706                .iter()
707                .zip(self.replay.index_operand_strides.iter())
708                .zip(self.spec.start_index_map.iter())
709            {
710                let index_offset = state.index_batch_offset + index_component_offset;
711                // SAFETY: checked index layout replay and RawStridedRef
712                // reachability prove this component read is in bounds.
713                let start = unsafe { *index_data.as_ptr().offset(index_offset) }.to_i64();
714                let clamped = self.clamp_window_start(start, operand_dim);
715                source_offset += operand_stride * clamped as isize;
716            }
717
718            // SAFETY: checked replay metadata and the validated writer prove
719            // both the source and the distinct destination offset are valid.
720            let value = unsafe { *operand_data.as_ptr().offset(source_offset) };
721            unsafe { dest.write_at(state.dest_offset, value) };
722            self.advance_replay_state(&mut state);
723        }
724        Ok(())
725    }
726
727    fn uses_rank_one_scalar_take_path(&self) -> bool {
728        self.operand_dims.len() == 1
729            && self.index_dims.len() == 1
730            && self.dest_dims.len() == 1
731            && self.operand_strides[0] == 1
732            && self.index_strides[0] == 1
733            && self.dest_strides[0] == 1
734            && self.spec.offset_dims.is_empty()
735            && self.spec.collapsed_slice_dims.as_slice() == [0]
736            && self.spec.start_index_map.as_slice() == [0]
737            && self.spec.index_vector_dim == 1
738            && self.spec.slice_sizes.as_slice() == [1]
739    }
740
741    fn execute_rank_one_scalar_take<T, I, W>(
742        &self,
743        dest: &mut W,
744        operand: &RawStridedRef<'_, T>,
745        start_indices: &RawStridedRef<'_, I>,
746    ) -> Result<()>
747    where
748        T: Copy,
749        I: GatherIndex,
750        W: OverwriteWriter<T>,
751    {
752        let mut dest_offset = dest.offset();
753        let mut index_offset = start_indices.offset();
754        let operand_offset = operand.offset();
755        let operand_data = operand.data();
756        let index_data = start_indices.data();
757
758        // INVARIANT: compile and check_call validated compact rank-one layouts,
759        // so incrementing these offsets once per logical element stays within
760        // the validated allocations and cannot overflow isize.
761        for _ in 0..self.total {
762            // SAFETY: the invariant above proves all three offsets are in bounds.
763            unsafe {
764                let start = (*index_data.as_ptr().offset(index_offset)).to_i64();
765                let source_offset = operand_offset + self.clamp_window_start(start, 0) as isize;
766                dest.write_at(dest_offset, *operand_data.as_ptr().offset(source_offset));
767            }
768            dest_offset += 1;
769            index_offset += 1;
770        }
771        Ok(())
772    }
773
774    #[cfg(feature = "parallel")]
775    fn execute_rank_one_scalar_take_parallel<T, I, W>(
776        &self,
777        dest: &mut W,
778        operand: &RawStridedRef<'_, T>,
779        start_indices: &RawStridedRef<'_, I>,
780        nthreads: usize,
781    ) -> Result<()>
782    where
783        T: Copy + MaybeSendSync,
784        I: GatherIndex,
785        W: OverwriteWriter<T>,
786    {
787        let dest_offset = dest.offset();
788        let operand_offset = operand.offset();
789        let index_offset = start_indices.offset();
790        // SAFETY: check_call validated the writer's complete destination layout.
791        let dest_ptr = crate::threading::SendPtr(unsafe { dest.data_ptr() });
792        let operand_ptr = crate::threading::SendPtr(operand.data().as_ptr() as *mut T);
793        let index_ptr = crate::threading::SendPtr(start_indices.data().as_ptr() as *mut I);
794
795        crate::threading::parallel_map_reduce(
796            0..self.total,
797            nthreads,
798            &|range| {
799                let dest_ptr = dest_ptr.as_ptr();
800                let operand_ptr = operand_ptr.as_const();
801                let index_ptr = index_ptr.as_const();
802                // INVARIANT: the validated compact rank-one layouts map each
803                // output position to one distinct destination offset.
804                for position in range {
805                    let position = position as isize;
806                    // SAFETY: the invariant above proves the source and distinct
807                    // destination offsets are in their validated allocations.
808                    unsafe {
809                        let start = (*index_ptr.offset(index_offset + position)).to_i64();
810                        let source_offset =
811                            operand_offset + self.clamp_window_start(start, 0) as isize;
812                        dest_ptr
813                            .offset(dest_offset + position)
814                            .write(operand_ptr.offset(source_offset).read());
815                    }
816                }
817                Ok(())
818            },
819            &|left, right| left.and(right),
820        )
821    }
822
823    #[cfg(feature = "parallel")]
824    fn execute_parallel<T, I, W>(
825        &self,
826        dest: &mut W,
827        operand: &RawStridedRef<'_, T>,
828        start_indices: &RawStridedRef<'_, I>,
829        nthreads: usize,
830    ) -> Result<()>
831    where
832        T: Copy + MaybeSendSync,
833        I: GatherIndex,
834        W: OverwriteWriter<T>,
835    {
836        let dest_offset_base = dest.offset();
837        let operand_offset_base = operand.offset();
838        let index_offset_base = start_indices.offset();
839        // SAFETY: the validated writer owns the destination allocation.
840        let dest_ptr = crate::threading::SendPtr(unsafe { dest.data_ptr() });
841        let operand_ptr = crate::threading::SendPtr(operand.data().as_ptr() as *mut T);
842        let index_ptr = crate::threading::SendPtr(start_indices.data().as_ptr() as *mut I);
843
844        crate::threading::parallel_map_reduce(
845            0..self.total,
846            nthreads,
847            &|range| {
848                let mut state = self.decode_replay_state(
849                    range.start,
850                    dest_offset_base,
851                    operand_offset_base,
852                    index_offset_base,
853                )?;
854                let dest_ptr = dest_ptr.as_ptr();
855                let operand_ptr = operand_ptr.as_const();
856                let index_ptr = index_ptr.as_const();
857
858                // INVARIANT: this worker owns a disjoint logical range. Checked
859                // range-start decode plus checked replay deltas keeps every
860                // state at its corresponding output, and injectivity makes all
861                // writes across workers distinct.
862                for _ in range {
863                    // INVARIANT: source_offset is freshly based on the current
864                    // window offset; clamped mapped coordinates remain inside
865                    // the validated operand span and are never accumulated.
866                    let mut source_offset = state.window_offset;
867                    for ((&index_component_offset, &operand_stride), &operand_dim) in self
868                        .replay
869                        .index_component_offsets
870                        .iter()
871                        .zip(self.replay.index_operand_strides.iter())
872                        .zip(self.spec.start_index_map.iter())
873                    {
874                        let index_offset = state.index_batch_offset + index_component_offset;
875                        // SAFETY: the checked index span and this worker's
876                        // validated range-start state prove this read is valid.
877                        let start = unsafe { (*index_ptr.offset(index_offset)).to_i64() };
878                        let clamped = self.clamp_window_start(start, operand_dim);
879                        source_offset += operand_stride * clamped as isize;
880                    }
881
882                    // SAFETY: the source proof above and distinct destination
883                    // invariant justify these raw accesses.
884                    unsafe {
885                        dest_ptr
886                            .offset(state.dest_offset)
887                            .write(operand_ptr.offset(source_offset).read());
888                    }
889                    self.advance_replay_state(&mut state);
890                }
891                Ok(())
892            },
893            &|left, right| left.and(right),
894        )
895    }
896
897    fn check_call<T, I, W>(
898        &self,
899        dest: &W,
900        operand: &RawStridedRef<'_, T>,
901        start_indices: &RawStridedRef<'_, I>,
902    ) -> Result<()>
903    where
904        W: OverwriteWriter<T>,
905    {
906        if dest.dims() != &self.dest_dims[..]
907            || dest.strides() != &self.dest_strides[..]
908            || operand.dims() != &self.operand_dims[..]
909            || operand.strides() != &self.operand_strides[..]
910            || start_indices.dims() != &self.index_dims[..]
911            || start_indices.strides() != &self.index_strides[..]
912        {
913            return Err(StridedError::PlanLayoutMismatch);
914        }
915        Ok(())
916    }
917
918    fn decode_replay_state(
919        &self,
920        mut linear: usize,
921        dest_offset_base: isize,
922        operand_offset_base: isize,
923        index_offset_base: isize,
924    ) -> Result<GatherReplayState> {
925        let mut coords = AxisVec::with_capacity(self.dest_dims.len());
926        let mut dest_offset = dest_offset_base;
927        let mut window_offset = operand_offset_base;
928        let mut index_batch_offset = index_offset_base;
929        for (axis, &dim) in self.dest_dims.iter().enumerate() {
930            let coord = linear % dim;
931            linear /= dim;
932            let replay_axis = self.replay.axes[axis];
933            coords.push(coord);
934            dest_offset = checked_offset_add(dest_offset, replay_axis.dest_step, coord)?;
935            window_offset = checked_offset_add(window_offset, replay_axis.window_step, coord)?;
936            index_batch_offset =
937                checked_offset_add(index_batch_offset, replay_axis.index_batch_step, coord)?;
938        }
939        Ok(GatherReplayState {
940            coords,
941            dest_offset,
942            window_offset,
943            index_batch_offset,
944        })
945    }
946
947    #[inline]
948    fn advance_replay_state(&self, state: &mut GatherReplayState) {
949        for (axis, replay_axis) in self.replay.axes.iter().enumerate() {
950            let coord = state.coords[axis];
951            if coord + 1 < self.dest_dims[axis] {
952                state.coords[axis] = coord + 1;
953                state.dest_offset += replay_axis.dest_step;
954                state.window_offset += replay_axis.window_step;
955                state.index_batch_offset += replay_axis.index_batch_step;
956                return;
957            }
958            state.coords[axis] = 0;
959            state.dest_offset += replay_axis.dest_reset;
960            state.window_offset += replay_axis.window_reset;
961            state.index_batch_offset += replay_axis.index_batch_reset;
962        }
963    }
964
965    #[inline]
966    fn clamp_window_start(&self, start: i64, operand_dim: usize) -> usize {
967        let dim_size = self.operand_dims[operand_dim];
968        let window_size = self.spec.slice_sizes[operand_dim];
969        let max_start = dim_size.saturating_sub(window_size) as i64;
970        start.clamp(0, max_start) as usize
971    }
972}
973
974impl DynamicSlicePlan {
975    /// Compile a fixed-window dynamic-slice traversal.
976    ///
977    /// `start_*` describes a rank-1 index vector whose length equals the
978    /// operand rank. `dest_dims` must equal `slice_sizes`.
979    pub fn compile(
980        operand_dims: &[usize],
981        operand_strides: &[isize],
982        start_dims: &[usize],
983        start_strides: &[isize],
984        dest_dims: &[usize],
985        dest_strides: &[isize],
986        slice_sizes: &[usize],
987    ) -> Result<Self> {
988        if operand_dims.len() != operand_strides.len()
989            || start_dims.len() != start_strides.len()
990            || dest_dims.len() != dest_strides.len()
991        {
992            return Err(StridedError::StrideLengthMismatch);
993        }
994        if slice_sizes.len() != operand_dims.len() {
995            return Err(StridedError::RankMismatch(
996                slice_sizes.len(),
997                operand_dims.len(),
998            ));
999        }
1000        validate_start_vector(start_dims, operand_dims.len())?;
1001        checked_total_len(operand_dims)?;
1002        checked_total_len(start_dims)?;
1003        let total = checked_total_len(dest_dims)?;
1004        validate_layout_span(operand_dims, operand_strides)?;
1005        validate_layout_span(start_dims, start_strides)?;
1006        validate_layout_span(dest_dims, dest_strides)?;
1007        if dest_dims != slice_sizes {
1008            return Err(StridedError::ShapeMismatch(
1009                dest_dims.to_vec(),
1010                slice_sizes.to_vec(),
1011            ));
1012        }
1013        if !crate::layout_check::is_injective_layout(dest_dims, dest_strides) {
1014            return Err(StridedError::NonInjectiveOutputLayout);
1015        }
1016        validate_window_sizes(operand_dims, slice_sizes)?;
1017        #[cfg(feature = "parallel")]
1018        WindowReplay::compile(slice_sizes, operand_strides, dest_strides)?;
1019        #[cfg(not(feature = "parallel"))]
1020        let replay = WindowReplay::compile(slice_sizes, operand_strides, dest_strides)?;
1021
1022        Ok(Self {
1023            operand_dims: operand_dims.into(),
1024            operand_strides: operand_strides.into(),
1025            start_dims: start_dims.into(),
1026            start_strides: start_strides.into(),
1027            dest_dims: dest_dims.into(),
1028            dest_strides: dest_strides.into(),
1029            slice_sizes: slice_sizes.into(),
1030            total,
1031            #[cfg(not(feature = "parallel"))]
1032            replay,
1033        })
1034    }
1035
1036    /// Execute the prepared dynamic-slice traversal.
1037    pub fn execute<T, I>(
1038        &self,
1039        dest: &mut RawStridedMut<'_, T>,
1040        operand: &RawStridedRef<'_, T>,
1041        starts: &RawStridedRef<'_, I>,
1042    ) -> Result<()>
1043    where
1044        T: Copy + MaybeSendSync,
1045        I: GatherIndex,
1046    {
1047        self.execute_with_writer(dest, operand, starts)
1048    }
1049
1050    pub(crate) fn execute_uninit<T, I>(
1051        &self,
1052        dest: &mut RawStridedMut<'_, MaybeUninit<T>>,
1053        operand: &RawStridedRef<'_, T>,
1054        starts: &RawStridedRef<'_, I>,
1055    ) -> Result<()>
1056    where
1057        T: Copy + MaybeSendSync,
1058        I: GatherIndex,
1059    {
1060        self.execute_with_writer(dest, operand, starts)
1061    }
1062
1063    fn execute_with_writer<T, I, W>(
1064        &self,
1065        dest: &mut W,
1066        operand: &RawStridedRef<'_, T>,
1067        starts: &RawStridedRef<'_, I>,
1068    ) -> Result<()>
1069    where
1070        T: Copy + MaybeSendSync,
1071        I: GatherIndex,
1072        W: OverwriteWriter<T>,
1073    {
1074        self.check_call(dest, operand, starts)?;
1075        if self.total == 0 {
1076            return Ok(());
1077        }
1078        if self.uses_rank_one_contiguous_path() {
1079            return self.execute_rank_one_contiguous(dest, operand, starts);
1080        }
1081        #[cfg(feature = "parallel")]
1082        let replay =
1083            WindowReplay::compile(&self.slice_sizes, &self.operand_strides, &self.dest_strides)?;
1084        #[cfg(not(feature = "parallel"))]
1085        let replay = &self.replay;
1086        #[cfg(feature = "parallel")]
1087        {
1088            let nthreads = crate::threading::parallel_threads_for_len(self.total);
1089            if nthreads > 1 {
1090                return self.execute_parallel(dest, operand, starts, nthreads, &replay);
1091            }
1092        }
1093
1094        let mut starts_storage = CoordScratch::new(self.operand_dims.len());
1095        let clamped_starts = starts_storage.as_mut_slice();
1096        read_clamped_starts(
1097            starts,
1098            &self.operand_dims,
1099            &self.slice_sizes,
1100            clamped_starts,
1101        )?;
1102        let source_base =
1103            checked_strided_offset(operand.offset(), &self.operand_strides, clamped_starts)?;
1104        let mut state = replay.decode(0, source_base, dest.offset())?;
1105        let operand_data = operand.data();
1106
1107        // INVARIANT: compile validated the full operand/destination spans and the
1108        // replay window spans; checked bases plus checked reset deltas therefore
1109        // keep every current offset reachable without per-element offset scans.
1110        for _ in 0..self.total {
1111            // SAFETY: the invariant above proves both current offsets are valid.
1112            let value = unsafe { *operand_data.as_ptr().offset(state.source_offset) };
1113            // SAFETY: the invariant above proves this logical destination offset
1114            // is in-bounds, and the output layout is injective.
1115            unsafe { dest.write_at(state.dest_offset, value) };
1116            replay.advance(&mut state);
1117        }
1118        Ok(())
1119    }
1120
1121    #[inline]
1122    fn uses_rank_one_contiguous_path(&self) -> bool {
1123        self.operand_dims.len() == 1 && self.operand_strides[0] == 1 && self.dest_strides[0] == 1
1124    }
1125
1126    fn execute_rank_one_contiguous<T, I, W>(
1127        &self,
1128        dest: &mut W,
1129        operand: &RawStridedRef<'_, T>,
1130        starts: &RawStridedRef<'_, I>,
1131    ) -> Result<()>
1132    where
1133        T: Copy,
1134        I: GatherIndex,
1135        W: OverwriteWriter<T>,
1136    {
1137        let mut clamped_starts = [0usize; 1];
1138        read_clamped_starts(
1139            starts,
1140            &self.operand_dims,
1141            &self.slice_sizes,
1142            &mut clamped_starts,
1143        )?;
1144        let source_start = checked_offset_add(operand.offset(), 1, clamped_starts[0])?;
1145        let source_start =
1146            usize::try_from(source_start).map_err(|_| StridedError::OffsetOverflow)?;
1147        let dest_start =
1148            usize::try_from(dest.offset()).map_err(|_| StridedError::OffsetOverflow)?;
1149        let source_end = source_start
1150            .checked_add(self.total)
1151            .ok_or(StridedError::OffsetOverflow)?;
1152        let source = operand
1153            .data()
1154            .get(source_start..source_end)
1155            .ok_or(StridedError::OffsetOverflow)?;
1156        // SAFETY: the validated writer owns the destination allocation.
1157        let dest_ptr = unsafe { dest.data_ptr() };
1158        // SAFETY: bounds were checked above and the writer owns the logical
1159        // destination storage.
1160        unsafe {
1161            core::ptr::copy_nonoverlapping(source.as_ptr(), dest_ptr.add(dest_start), self.total);
1162        }
1163        Ok(())
1164    }
1165
1166    #[cfg(feature = "parallel")]
1167    fn execute_parallel<T, I, W>(
1168        &self,
1169        dest: &mut W,
1170        operand: &RawStridedRef<'_, T>,
1171        starts: &RawStridedRef<'_, I>,
1172        nthreads: usize,
1173        replay: &WindowReplay,
1174    ) -> Result<()>
1175    where
1176        T: Copy + MaybeSendSync,
1177        I: GatherIndex,
1178        W: OverwriteWriter<T>,
1179    {
1180        let mut clamped_starts: AxisVec<usize> = (0..self.operand_dims.len()).map(|_| 0).collect();
1181        read_clamped_starts(
1182            starts,
1183            &self.operand_dims,
1184            &self.slice_sizes,
1185            &mut clamped_starts,
1186        )?;
1187        let source_base =
1188            checked_strided_offset(operand.offset(), &self.operand_strides, &clamped_starts)?;
1189        let dest_base = dest.offset();
1190        let operand_ptr = crate::threading::SendPtr(operand.data().as_ptr() as *mut T);
1191        // SAFETY: the validated writer owns the destination allocation.
1192        let dest_ptr = crate::threading::SendPtr(unsafe { dest.data_ptr() });
1193
1194        crate::threading::parallel_map_reduce(
1195            0..self.total,
1196            nthreads,
1197            &|range| {
1198                let mut state = replay.decode(range.start, source_base, dest_base)?;
1199                let operand_ptr = operand_ptr.as_const();
1200                let dest_ptr = dest_ptr.as_ptr();
1201
1202                // INVARIANT: each worker decodes one checked range start. Its
1203                // replay range is disjoint, and the injective output layout makes
1204                // all writes distinct across workers.
1205                for _ in range {
1206                    // SAFETY: checked replay bases/resets keep both offsets in
1207                    // their validated allocations.
1208                    unsafe {
1209                        dest_ptr
1210                            .offset(state.dest_offset)
1211                            .write(operand_ptr.offset(state.source_offset).read());
1212                    }
1213                    replay.advance(&mut state);
1214                }
1215                Ok(())
1216            },
1217            &|left, right| left.and(right),
1218        )
1219    }
1220
1221    fn check_call<T, I, W>(
1222        &self,
1223        dest: &W,
1224        operand: &RawStridedRef<'_, T>,
1225        starts: &RawStridedRef<'_, I>,
1226    ) -> Result<()>
1227    where
1228        W: OverwriteWriter<T>,
1229    {
1230        if dest.dims() != &self.dest_dims[..]
1231            || dest.strides() != &self.dest_strides[..]
1232            || operand.dims() != &self.operand_dims[..]
1233            || operand.strides() != &self.operand_strides[..]
1234            || starts.dims() != &self.start_dims[..]
1235            || starts.strides() != &self.start_strides[..]
1236        {
1237            return Err(StridedError::PlanLayoutMismatch);
1238        }
1239        Ok(())
1240    }
1241}
1242
1243impl DynamicUpdateSlicePlan {
1244    /// Compile a dynamic-update-slice traversal.
1245    ///
1246    /// `dest_dims` must match `operand_dims`; execution materializes
1247    /// `dest = operand` and then overwrites the clamped update window.
1248    #[allow(clippy::too_many_arguments)]
1249    pub fn compile(
1250        operand_dims: &[usize],
1251        operand_strides: &[isize],
1252        start_dims: &[usize],
1253        start_strides: &[isize],
1254        update_dims: &[usize],
1255        update_strides: &[isize],
1256        dest_dims: &[usize],
1257        dest_strides: &[isize],
1258    ) -> Result<Self> {
1259        if operand_dims.len() != operand_strides.len()
1260            || start_dims.len() != start_strides.len()
1261            || update_dims.len() != update_strides.len()
1262            || dest_dims.len() != dest_strides.len()
1263        {
1264            return Err(StridedError::StrideLengthMismatch);
1265        }
1266        if update_dims.len() != operand_dims.len() {
1267            return Err(StridedError::RankMismatch(
1268                update_dims.len(),
1269                operand_dims.len(),
1270            ));
1271        }
1272        validate_start_vector(start_dims, operand_dims.len())?;
1273        checked_total_len(operand_dims)?;
1274        checked_total_len(start_dims)?;
1275        let total = checked_total_len(update_dims)?;
1276        validate_layout_span(operand_dims, operand_strides)?;
1277        validate_layout_span(start_dims, start_strides)?;
1278        validate_layout_span(update_dims, update_strides)?;
1279        validate_layout_span(dest_dims, dest_strides)?;
1280        if dest_dims != operand_dims {
1281            return Err(StridedError::ShapeMismatch(
1282                dest_dims.to_vec(),
1283                operand_dims.to_vec(),
1284            ));
1285        }
1286        if !crate::layout_check::is_injective_layout(dest_dims, dest_strides) {
1287            return Err(StridedError::NonInjectiveOutputLayout);
1288        }
1289        validate_window_sizes(operand_dims, update_dims)?;
1290        let copy_plan = CopyPlan::compile(operand_dims, dest_strides, operand_strides)?;
1291        #[cfg(feature = "parallel")]
1292        WindowReplay::compile(update_dims, update_strides, dest_strides)?;
1293        #[cfg(not(feature = "parallel"))]
1294        let replay = WindowReplay::compile(update_dims, update_strides, dest_strides)?;
1295
1296        Ok(Self {
1297            operand_dims: operand_dims.into(),
1298            operand_strides: operand_strides.into(),
1299            start_dims: start_dims.into(),
1300            start_strides: start_strides.into(),
1301            update_dims: update_dims.into(),
1302            update_strides: update_strides.into(),
1303            dest_dims: dest_dims.into(),
1304            dest_strides: dest_strides.into(),
1305            total,
1306            copy_plan,
1307            #[cfg(not(feature = "parallel"))]
1308            replay,
1309        })
1310    }
1311
1312    /// Execute the prepared dynamic-update-slice traversal.
1313    pub fn execute<T, I>(
1314        &self,
1315        dest: &mut RawStridedMut<'_, T>,
1316        operand: &RawStridedRef<'_, T>,
1317        update: &RawStridedRef<'_, T>,
1318        starts: &RawStridedRef<'_, I>,
1319    ) -> Result<()>
1320    where
1321        T: Copy + MaybeSendSync,
1322        I: GatherIndex,
1323    {
1324        self.check_call(dest, operand, update, starts)?;
1325        self.copy_plan.execute(dest, operand)?;
1326        self.execute_update_with_writer(dest, update, starts)
1327    }
1328
1329    /// Execute dynamic update into a destination whose reachable slots may be
1330    /// uninitialized. The copy completes before any read-modify-write access.
1331    pub(crate) fn execute_uninit<'a, T, I>(
1332        &self,
1333        dest: &'a mut RawStridedMut<'a, MaybeUninit<T>>,
1334        operand: &RawStridedRef<'_, T>,
1335        update: &RawStridedRef<'_, T>,
1336        starts: &RawStridedRef<'_, I>,
1337    ) -> Result<()>
1338    where
1339        T: Copy + MaybeSendSync,
1340        I: GatherIndex,
1341    {
1342        self.check_call(dest, operand, update, starts)?;
1343        self.copy_plan
1344            .execute_uninit_then(dest, operand, |mut receipt| {
1345                self.execute_update_with_writer(&mut receipt, update, starts)
1346            })?
1347    }
1348
1349    fn execute_update_with_writer<T, I, W>(
1350        &self,
1351        dest: &mut W,
1352        update: &RawStridedRef<'_, T>,
1353        starts: &RawStridedRef<'_, I>,
1354    ) -> Result<()>
1355    where
1356        T: Copy + MaybeSendSync,
1357        I: GatherIndex,
1358        W: OverwriteWriter<T>,
1359    {
1360        if self.total == 0 {
1361            return Ok(());
1362        }
1363        if self.uses_rank_one_contiguous_path() {
1364            return self.execute_rank_one_contiguous(dest, update, starts);
1365        }
1366        #[cfg(feature = "parallel")]
1367        let replay =
1368            WindowReplay::compile(&self.update_dims, &self.update_strides, &self.dest_strides)?;
1369        #[cfg(not(feature = "parallel"))]
1370        let replay = &self.replay;
1371        #[cfg(feature = "parallel")]
1372        {
1373            let nthreads = crate::threading::parallel_threads_for_len(self.total);
1374            if nthreads > 1 {
1375                return self.execute_update_parallel(dest, update, starts, nthreads, &replay);
1376            }
1377        }
1378
1379        let mut starts_storage = CoordScratch::new(self.operand_dims.len());
1380        let clamped_starts = starts_storage.as_mut_slice();
1381        read_clamped_starts(
1382            starts,
1383            &self.operand_dims,
1384            &self.update_dims,
1385            clamped_starts,
1386        )?;
1387        let source_base = update.offset();
1388        let dest_base = checked_strided_offset(dest.offset(), &self.dest_strides, &clamped_starts)?;
1389        let mut state = replay.decode(0, source_base, dest_base)?;
1390        let update_data = update.data();
1391
1392        // INVARIANT: the initial CopyPlan has completed before this replay;
1393        // compile validated both window spans and the checked bases/resets keep
1394        // every update read and destination write reachable.
1395        for _ in 0..self.total {
1396            // SAFETY: checked replay metadata proves the current update read.
1397            let value = unsafe { *update_data.as_ptr().offset(state.source_offset) };
1398            // SAFETY: the copied, injective destination layout proves this write
1399            // is initialized and in-bounds.
1400            unsafe { dest.write_at(state.dest_offset, value) };
1401            replay.advance(&mut state);
1402        }
1403        Ok(())
1404    }
1405
1406    #[inline]
1407    fn uses_rank_one_contiguous_path(&self) -> bool {
1408        self.operand_dims.len() == 1
1409            && self.operand_strides[0] == 1
1410            && self.update_strides[0] == 1
1411            && self.dest_strides[0] == 1
1412    }
1413
1414    fn execute_rank_one_contiguous<T, I, W>(
1415        &self,
1416        dest: &mut W,
1417        update: &RawStridedRef<'_, T>,
1418        starts: &RawStridedRef<'_, I>,
1419    ) -> Result<()>
1420    where
1421        T: Copy,
1422        I: GatherIndex,
1423        W: OverwriteWriter<T>,
1424    {
1425        let mut clamped_starts = [0usize; 1];
1426        read_clamped_starts(
1427            starts,
1428            &self.operand_dims,
1429            &self.update_dims,
1430            &mut clamped_starts,
1431        )?;
1432        let update_start =
1433            usize::try_from(update.offset()).map_err(|_| StridedError::OffsetOverflow)?;
1434        let dest_start = checked_offset_add(dest.offset(), 1, clamped_starts[0])?;
1435        let dest_start = usize::try_from(dest_start).map_err(|_| StridedError::OffsetOverflow)?;
1436        let update_end = update_start
1437            .checked_add(self.total)
1438            .ok_or(StridedError::OffsetOverflow)?;
1439        let update = update
1440            .data()
1441            .get(update_start..update_end)
1442            .ok_or(StridedError::OffsetOverflow)?;
1443        // SAFETY: the validated writer owns the destination allocation.
1444        let dest_ptr = unsafe { dest.data_ptr() };
1445        // SAFETY: the checked ranges are inside the destination allocation.
1446        unsafe {
1447            core::ptr::copy_nonoverlapping(update.as_ptr(), dest_ptr.add(dest_start), self.total);
1448        }
1449        Ok(())
1450    }
1451
1452    #[cfg(feature = "parallel")]
1453    fn execute_update_parallel<T, I, W>(
1454        &self,
1455        dest: &mut W,
1456        update: &RawStridedRef<'_, T>,
1457        starts: &RawStridedRef<'_, I>,
1458        nthreads: usize,
1459        replay: &WindowReplay,
1460    ) -> Result<()>
1461    where
1462        T: Copy + MaybeSendSync,
1463        I: GatherIndex,
1464        W: OverwriteWriter<T>,
1465    {
1466        let mut clamped_starts: AxisVec<usize> = (0..self.operand_dims.len()).map(|_| 0).collect();
1467        read_clamped_starts(
1468            starts,
1469            &self.operand_dims,
1470            &self.update_dims,
1471            &mut clamped_starts,
1472        )?;
1473        let source_base = update.offset();
1474        let dest_base = checked_strided_offset(dest.offset(), &self.dest_strides, &clamped_starts)?;
1475        let update_ptr = crate::threading::SendPtr(update.data().as_ptr() as *mut T);
1476        // SAFETY: the validated writer owns the destination allocation.
1477        let dest_ptr = crate::threading::SendPtr(unsafe { dest.data_ptr() });
1478
1479        crate::threading::parallel_map_reduce(
1480            0..self.total,
1481            nthreads,
1482            &|range| {
1483                let mut state = replay.decode(range.start, source_base, dest_base)?;
1484                let update_ptr = update_ptr.as_const();
1485                let dest_ptr = dest_ptr.as_ptr();
1486
1487                // INVARIANT: the initial operand copy completed before this
1488                // replay; workers own disjoint ranges and injective output
1489                // layouts make their writes distinct.
1490                for _ in range {
1491                    // SAFETY: checked replay bases/resets prove both offsets are
1492                    // within the initialized update and destination allocations.
1493                    unsafe {
1494                        dest_ptr
1495                            .offset(state.dest_offset)
1496                            .write(update_ptr.offset(state.source_offset).read());
1497                    }
1498                    replay.advance(&mut state);
1499                }
1500                Ok(())
1501            },
1502            &|left, right| left.and(right),
1503        )
1504    }
1505
1506    fn check_call<T, I, W>(
1507        &self,
1508        dest: &W,
1509        operand: &RawStridedRef<'_, T>,
1510        update: &RawStridedRef<'_, T>,
1511        starts: &RawStridedRef<'_, I>,
1512    ) -> Result<()>
1513    where
1514        W: OverwriteWriter<T>,
1515    {
1516        if dest.dims() != &self.dest_dims[..]
1517            || dest.strides() != &self.dest_strides[..]
1518            || operand.dims() != &self.operand_dims[..]
1519            || operand.strides() != &self.operand_strides[..]
1520            || update.dims() != &self.update_dims[..]
1521            || update.strides() != &self.update_strides[..]
1522            || starts.dims() != &self.start_dims[..]
1523            || starts.strides() != &self.start_strides[..]
1524        {
1525            return Err(StridedError::PlanLayoutMismatch);
1526        }
1527        Ok(())
1528    }
1529}
1530
1531impl ScatterPlan {
1532    /// Compile an additive scatter traversal.
1533    #[allow(clippy::too_many_arguments)]
1534    pub fn compile(
1535        operand_dims: &[usize],
1536        operand_strides: &[isize],
1537        index_dims: &[usize],
1538        index_strides: &[isize],
1539        update_dims: &[usize],
1540        update_strides: &[isize],
1541        dest_dims: &[usize],
1542        dest_strides: &[isize],
1543        spec: ScatterSpec,
1544    ) -> Result<Self> {
1545        if operand_dims.len() != operand_strides.len()
1546            || index_dims.len() != index_strides.len()
1547            || update_dims.len() != update_strides.len()
1548            || dest_dims.len() != dest_strides.len()
1549        {
1550            return Err(StridedError::StrideLengthMismatch);
1551        }
1552        checked_total_len(operand_dims)?;
1553        checked_total_len(index_dims)?;
1554        checked_total_len(update_dims)?;
1555        if dest_dims != operand_dims {
1556            return Err(StridedError::ShapeMismatch(
1557                dest_dims.to_vec(),
1558                operand_dims.to_vec(),
1559            ));
1560        }
1561        if !crate::layout_check::is_injective_layout(dest_dims, dest_strides) {
1562            return Err(StridedError::NonInjectiveOutputLayout);
1563        }
1564
1565        let operand_rank = operand_dims.len();
1566        validate_unique_axes(&spec.inserted_window_dims, operand_rank)?;
1567        validate_unique_axes(&spec.scatter_dims_to_operand_dims, operand_rank)?;
1568        if spec.index_vector_dim > index_dims.len() {
1569            return Err(StridedError::InvalidAxis {
1570                axis: spec.index_vector_dim,
1571                rank: index_dims.len() + 1,
1572            });
1573        }
1574        let index_vector_size = if spec.index_vector_dim == index_dims.len() {
1575            1
1576        } else {
1577            index_dims[spec.index_vector_dim]
1578        };
1579        if index_vector_size != spec.scatter_dims_to_operand_dims.len() {
1580            return Err(StridedError::RankMismatch(
1581                index_vector_size,
1582                spec.scatter_dims_to_operand_dims.len(),
1583            ));
1584        }
1585
1586        let batch_shape = index_batch_shape(index_dims, spec.index_vector_dim);
1587        let window_dims = operand_window_dims(operand_rank, &spec.inserted_window_dims);
1588        if spec.update_window_dims.len() != window_dims.len() {
1589            return Err(StridedError::RankMismatch(
1590                spec.update_window_dims.len(),
1591                window_dims.len(),
1592            ));
1593        }
1594
1595        let update_rank = update_dims.len();
1596        let expected_batch_rank = update_rank
1597            .checked_sub(spec.update_window_dims.len())
1598            .ok_or(StridedError::RankMismatch(
1599                spec.update_window_dims.len(),
1600                update_rank,
1601            ))?;
1602        if expected_batch_rank != batch_shape.len() {
1603            return Err(StridedError::RankMismatch(
1604                expected_batch_rank,
1605                batch_shape.len(),
1606            ));
1607        }
1608        validate_unique_axes(&spec.update_window_dims, update_rank)?;
1609
1610        let mut is_update_window_dim: AxisVec<bool> = (0..update_rank).map(|_| false).collect();
1611        for &axis in &spec.update_window_dims {
1612            is_update_window_dim[axis] = true;
1613        }
1614
1615        let mut batch_axis = 0usize;
1616        for axis in 0..update_rank {
1617            if !is_update_window_dim[axis] {
1618                if update_dims[axis] != batch_shape[batch_axis] {
1619                    return Err(StridedError::ShapeMismatch(
1620                        update_dims.to_vec(),
1621                        expected_scatter_update_shape(&batch_shape, &spec, update_dims).to_vec(),
1622                    ));
1623                }
1624                batch_axis += 1;
1625            }
1626        }
1627
1628        let mut window_shape: AxisVec<usize> = (0..operand_rank).map(|_| 1).collect();
1629        let mut window_shape_updates: AxisVec<usize> =
1630            AxisVec::with_capacity(spec.update_window_dims.len());
1631        for (pos, &update_axis) in spec.update_window_dims.iter().enumerate() {
1632            let dim = update_dims[update_axis];
1633            window_shape_updates.push(dim);
1634            window_shape[window_dims[pos]] = dim;
1635        }
1636        validate_window_sizes(operand_dims, &window_shape)?;
1637
1638        let batch_elems = checked_total_len(&batch_shape)?;
1639        let window_elems = checked_total_len(&window_shape_updates)?;
1640        let copy_plan = CopyPlan::compile(operand_dims, dest_strides, operand_strides)?;
1641        #[cfg(feature = "parallel")]
1642        ScatterReplay::validate(
1643            &batch_shape,
1644            index_dims,
1645            index_strides,
1646            spec.index_vector_dim,
1647            spec.scatter_dims_to_operand_dims.len(),
1648            update_dims,
1649            update_strides,
1650            &spec.update_window_dims,
1651            &is_update_window_dim,
1652            &window_shape_updates,
1653            &window_dims,
1654            dest_strides,
1655        )?;
1656        #[cfg(not(feature = "parallel"))]
1657        let replay = ScatterReplay::compile(
1658            &batch_shape,
1659            index_dims,
1660            index_strides,
1661            spec.index_vector_dim,
1662            spec.scatter_dims_to_operand_dims.len(),
1663            update_dims,
1664            update_strides,
1665            &spec.update_window_dims,
1666            &is_update_window_dim,
1667            &window_shape_updates,
1668            &window_dims,
1669            dest_strides,
1670        )?;
1671
1672        Ok(Self {
1673            operand_dims: operand_dims.into(),
1674            operand_strides: operand_strides.into(),
1675            index_dims: index_dims.into(),
1676            index_strides: index_strides.into(),
1677            update_dims: update_dims.into(),
1678            update_strides: update_strides.into(),
1679            dest_dims: dest_dims.into(),
1680            dest_strides: dest_strides.into(),
1681            spec,
1682            #[cfg(feature = "parallel")]
1683            batch_shape,
1684            #[cfg(feature = "parallel")]
1685            window_dims,
1686            window_shape,
1687            #[cfg(feature = "parallel")]
1688            window_shape_updates,
1689            #[cfg(feature = "parallel")]
1690            is_update_window_dim,
1691            batch_elems,
1692            window_elems,
1693            copy_plan,
1694            #[cfg(not(feature = "parallel"))]
1695            replay,
1696        })
1697    }
1698
1699    /// Execute the prepared additive scatter traversal.
1700    #[inline(always)]
1701    pub fn execute<T, I>(
1702        &self,
1703        dest: &mut RawStridedMut<'_, T>,
1704        operand: &RawStridedRef<'_, T>,
1705        scatter_indices: &RawStridedRef<'_, I>,
1706        updates: &RawStridedRef<'_, T>,
1707    ) -> Result<()>
1708    where
1709        T: Copy + Add<Output = T> + MaybeSendSync,
1710        I: GatherIndex,
1711    {
1712        self.check_call(dest, operand, scatter_indices, updates)?;
1713        self.copy_plan.execute(dest, operand)?;
1714        self.execute_updates(dest, scatter_indices, updates, |a, b| a + b)
1715    }
1716
1717    /// Execute additive scatter into a destination whose reachable slots may
1718    /// be uninitialized. The operand copy completes before any RMW access.
1719    pub(crate) fn execute_uninit<'a, T, I>(
1720        &self,
1721        dest: &'a mut RawStridedMut<'a, MaybeUninit<T>>,
1722        operand: &RawStridedRef<'_, T>,
1723        scatter_indices: &RawStridedRef<'_, I>,
1724        updates: &RawStridedRef<'_, T>,
1725        combine: fn(T, T) -> T,
1726    ) -> Result<()>
1727    where
1728        T: Copy + Add<Output = T> + MaybeSendSync,
1729        I: GatherIndex,
1730    {
1731        self.check_call(dest, operand, scatter_indices, updates)?;
1732        self.copy_plan
1733            .execute_uninit_then(dest, operand, |mut receipt| {
1734                self.execute_updates(&mut receipt, scatter_indices, updates, combine)
1735            })?
1736    }
1737
1738    #[inline(always)]
1739    fn execute_updates<T, I, W, F>(
1740        &self,
1741        dest: &mut W,
1742        scatter_indices: &RawStridedRef<'_, I>,
1743        updates: &RawStridedRef<'_, T>,
1744        combine: F,
1745    ) -> Result<()>
1746    where
1747        T: Copy + MaybeSendSync,
1748        I: GatherIndex,
1749        W: ReadModifyWrite<T>,
1750        F: Fn(T, T) -> T + Copy,
1751    {
1752        if self.batch_elems == 0 || self.window_elems == 0 {
1753            return Ok(());
1754        }
1755        if self.uses_rank_one_scalar_update_path() {
1756            return self.execute_rank_one_scalar_updates(dest, scatter_indices, updates, combine);
1757        }
1758        self.execute_generic_updates(dest, scatter_indices, updates, combine)
1759    }
1760
1761    #[inline(never)]
1762    fn execute_generic_updates<T, I, W, F>(
1763        &self,
1764        dest: &mut W,
1765        scatter_indices: &RawStridedRef<'_, I>,
1766        updates: &RawStridedRef<'_, T>,
1767        combine: F,
1768    ) -> Result<()>
1769    where
1770        T: Copy + MaybeSendSync,
1771        I: GatherIndex,
1772        W: ReadModifyWrite<T>,
1773        F: Fn(T, T) -> T + Copy,
1774    {
1775        // Overlapping additive updates are order-sensitive, so this remains a
1776        // deterministic serial replay until a combine-aware parallel plan exists.
1777        #[cfg(feature = "parallel")]
1778        let replay = ScatterReplay::compile(
1779            &self.batch_shape,
1780            &self.index_dims,
1781            &self.index_strides,
1782            self.spec.index_vector_dim,
1783            self.spec.scatter_dims_to_operand_dims.len(),
1784            &self.update_dims,
1785            &self.update_strides,
1786            &self.spec.update_window_dims,
1787            &self.is_update_window_dim,
1788            &self.window_shape_updates,
1789            &self.window_dims,
1790            &self.dest_strides,
1791        )?;
1792        #[cfg(not(feature = "parallel"))]
1793        let replay = &self.replay;
1794
1795        let mut operand_base_storage = CoordScratch::new(self.operand_dims.len());
1796        let operand_base = operand_base_storage.as_mut_slice();
1797        let mut batch_state = replay
1798            .batch
1799            .decode(0, scatter_indices.offset(), updates.offset())?;
1800        let index_data = scatter_indices.data();
1801        let update_data = updates.data();
1802
1803        // INVARIANT: batch and window replay metadata was checked at compile;
1804        // the indirect component loop below is the only data-dependent lookup.
1805        for _ in 0..self.batch_elems {
1806            operand_base.fill(0);
1807            for (&component_offset, &operand_axis) in replay
1808                .index_component_offsets
1809                .iter()
1810                .zip(self.spec.scatter_dims_to_operand_dims.iter())
1811            {
1812                let index_offset = batch_state
1813                    .source_offset
1814                    .checked_add(component_offset)
1815                    .ok_or(StridedError::OffsetOverflow)?;
1816                // SAFETY: checked batch replay and component offsets keep this
1817                // indirect index read inside the validated index allocation.
1818                let start = unsafe { *index_data.as_ptr().offset(index_offset) }.to_i64();
1819                operand_base[operand_axis] = clamp_window_start(
1820                    start,
1821                    self.operand_dims[operand_axis],
1822                    self.window_shape[operand_axis],
1823                );
1824            }
1825
1826            let dest_base = checked_strided_offset(dest.offset(), dest.strides(), operand_base)?;
1827            let mut window_state = replay
1828                .window
1829                .decode(0, batch_state.dest_offset, dest_base)?;
1830            for _ in 0..self.window_elems {
1831                // SAFETY: copy completion and checked replay state prove both
1832                // the update read and initialized destination RMW are valid.
1833                let value = unsafe { *update_data.as_ptr().offset(window_state.source_offset) };
1834                unsafe { dest.add_at(window_state.dest_offset, value, combine) };
1835                replay.window.advance(&mut window_state);
1836            }
1837            replay.batch.advance(&mut batch_state);
1838        }
1839        Ok(())
1840    }
1841
1842    fn uses_rank_one_scalar_update_path(&self) -> bool {
1843        self.operand_dims.len() == 1
1844            && self.index_dims.len() == 2
1845            && self.update_dims.len() == 1
1846            && self.dest_dims.len() == 1
1847            && self.operand_strides[0] == 1
1848            && self.index_strides[0] == 1
1849            && self.update_strides[0] == 1
1850            && self.dest_strides[0] == 1
1851            && self.index_dims[1] == 1
1852            && self.spec.update_window_dims.is_empty()
1853            && self.spec.inserted_window_dims.as_slice() == [0]
1854            && self.spec.scatter_dims_to_operand_dims.as_slice() == [0]
1855            && self.spec.index_vector_dim == 1
1856            && self.window_elems == 1
1857    }
1858
1859    #[inline(always)]
1860    fn execute_rank_one_scalar_updates<T, I, W, F>(
1861        &self,
1862        dest: &mut W,
1863        scatter_indices: &RawStridedRef<'_, I>,
1864        updates: &RawStridedRef<'_, T>,
1865        combine: F,
1866    ) -> Result<()>
1867    where
1868        T: Copy,
1869        I: GatherIndex,
1870        W: ReadModifyWrite<T>,
1871        F: Fn(T, T) -> T + Copy,
1872    {
1873        let mut index_offset = scatter_indices.offset();
1874        let mut update_offset = updates.offset();
1875        let dest_offset = dest.offset();
1876        let index_data = scatter_indices.data();
1877        let update_data = updates.data();
1878        // SAFETY: the validated writer owns the destination allocation.
1879        let dest_ptr = unsafe { dest.data_ptr() };
1880
1881        // INVARIANT: compile and check_call validated compact rank-one index
1882        // and update layouts. Ordered replay preserves repeated-index semantics.
1883        for _ in 0..self.batch_elems {
1884            // SAFETY: the invariant above proves index/update reads and the
1885            // clamped destination offset are in their validated allocations.
1886            unsafe {
1887                let start = (*index_data.as_ptr().offset(index_offset)).to_i64();
1888                let output_offset =
1889                    dest_offset + clamp_window_start(start, self.operand_dims[0], 1) as isize;
1890                let output = dest_ptr.offset(output_offset);
1891                output.write(combine(
1892                    output.read(),
1893                    *update_data.as_ptr().offset(update_offset),
1894                ));
1895            }
1896            index_offset += 1;
1897            update_offset += 1;
1898        }
1899        Ok(())
1900    }
1901
1902    fn check_call<T, I, W>(
1903        &self,
1904        dest: &W,
1905        operand: &RawStridedRef<'_, T>,
1906        scatter_indices: &RawStridedRef<'_, I>,
1907        updates: &RawStridedRef<'_, T>,
1908    ) -> Result<()>
1909    where
1910        W: OverwriteWriter<T>,
1911    {
1912        if dest.dims() != &self.dest_dims[..]
1913            || dest.strides() != &self.dest_strides[..]
1914            || operand.dims() != &self.operand_dims[..]
1915            || operand.strides() != &self.operand_strides[..]
1916            || scatter_indices.dims() != &self.index_dims[..]
1917            || scatter_indices.strides() != &self.index_strides[..]
1918            || updates.dims() != &self.update_dims[..]
1919            || updates.strides() != &self.update_strides[..]
1920        {
1921            return Err(StridedError::PlanLayoutMismatch);
1922        }
1923        Ok(())
1924    }
1925}
1926
1927fn validate_unique_axes(axes: &[usize], rank: usize) -> Result<()> {
1928    let mut seen = vec![false; rank];
1929    for &axis in axes {
1930        if axis >= rank {
1931            return Err(StridedError::InvalidAxis { axis, rank });
1932        }
1933        if seen[axis] {
1934            return Err(StridedError::InvalidAxis { axis, rank });
1935        }
1936        seen[axis] = true;
1937    }
1938    Ok(())
1939}
1940
1941fn validate_start_vector(start_dims: &[usize], operand_rank: usize) -> Result<()> {
1942    if start_dims.len() != 1 {
1943        return Err(StridedError::RankMismatch(start_dims.len(), 1));
1944    }
1945    if start_dims[0] != operand_rank {
1946        return Err(StridedError::RankMismatch(start_dims[0], operand_rank));
1947    }
1948    Ok(())
1949}
1950
1951fn validate_window_sizes(operand_dims: &[usize], window_sizes: &[usize]) -> Result<()> {
1952    if operand_dims.len() != window_sizes.len() {
1953        return Err(StridedError::RankMismatch(
1954            window_sizes.len(),
1955            operand_dims.len(),
1956        ));
1957    }
1958    for (axis, (&window, &dim)) in window_sizes.iter().zip(operand_dims.iter()).enumerate() {
1959        if window > dim {
1960            return Err(StridedError::InvalidAxis {
1961                axis,
1962                rank: operand_dims.len(),
1963            });
1964        }
1965    }
1966    Ok(())
1967}
1968
1969fn read_clamped_starts<I>(
1970    starts: &RawStridedRef<'_, I>,
1971    operand_dims: &[usize],
1972    window_sizes: &[usize],
1973    out: &mut [usize],
1974) -> Result<()>
1975where
1976    I: GatherIndex,
1977{
1978    debug_assert_eq!(operand_dims.len(), window_sizes.len());
1979    debug_assert_eq!(operand_dims.len(), out.len());
1980    for axis in 0..operand_dims.len() {
1981        let offset = checked_offset_add(starts.offset(), starts.strides()[0], axis)?;
1982        let start = unsafe { *starts.data().as_ptr().offset(offset) }.to_i64();
1983        out[axis] = clamp_window_start(start, operand_dims[axis], window_sizes[axis]);
1984    }
1985    Ok(())
1986}
1987
1988#[inline]
1989fn clamp_window_start(start: i64, dim_size: usize, window_size: usize) -> usize {
1990    let max_start = dim_size.saturating_sub(window_size) as i64;
1991    start.clamp(0, max_start) as usize
1992}
1993
1994fn expected_scatter_update_shape(
1995    batch_shape: &[usize],
1996    spec: &ScatterSpec,
1997    update_dims: &[usize],
1998) -> AxisVec<usize> {
1999    let mut expected: AxisVec<usize> = AxisVec::with_capacity(update_dims.len());
2000    let mut batch_axis = 0usize;
2001    for axis in 0..update_dims.len() {
2002        if spec.update_window_dims.contains(&axis) {
2003            expected.push(update_dims[axis]);
2004        } else {
2005            expected.push(batch_shape[batch_axis]);
2006            batch_axis += 1;
2007        }
2008    }
2009    expected
2010}
2011
2012fn operand_window_dims(rank: usize, collapsed_slice_dims: &[usize]) -> AxisVec<usize> {
2013    (0..rank)
2014        .filter(|axis| !collapsed_slice_dims.contains(axis))
2015        .collect()
2016}
2017
2018fn index_batch_shape(index_dims: &[usize], index_vector_dim: usize) -> AxisVec<usize> {
2019    if index_vector_dim == index_dims.len() {
2020        return index_dims.into();
2021    }
2022    index_dims
2023        .iter()
2024        .enumerate()
2025        .filter_map(|(axis, &dim)| (axis != index_vector_dim).then_some(dim))
2026        .collect()
2027}
2028
2029fn validate_layout_span(dims: &[usize], strides: &[isize]) -> Result<()> {
2030    if dims.len() != strides.len() {
2031        return Err(StridedError::StrideLengthMismatch);
2032    }
2033    let mut min_offset = 0isize;
2034    let mut max_offset = 0isize;
2035    for (&dim, &stride) in dims.iter().zip(strides.iter()) {
2036        let last =
2037            isize::try_from(dim.saturating_sub(1)).map_err(|_| StridedError::OffsetOverflow)?;
2038        let extent = stride
2039            .checked_mul(last)
2040            .ok_or(StridedError::OffsetOverflow)?;
2041        if extent < 0 {
2042            min_offset = min_offset
2043                .checked_add(extent)
2044                .ok_or(StridedError::OffsetOverflow)?;
2045        } else {
2046            max_offset = max_offset
2047                .checked_add(extent)
2048                .ok_or(StridedError::OffsetOverflow)?;
2049        }
2050    }
2051    Ok(())
2052}
2053
2054fn checked_replay_reset(dim: usize, step: isize) -> Result<isize> {
2055    if dim == 0 {
2056        return Ok(0);
2057    }
2058    let last = isize::try_from(dim - 1).map_err(|_| StridedError::OffsetOverflow)?;
2059    step.checked_mul(last)
2060        .and_then(isize::checked_neg)
2061        .ok_or(StridedError::OffsetOverflow)
2062}
2063
2064fn checked_total_len(dims: &[usize]) -> Result<usize> {
2065    if dims.is_empty() {
2066        return Ok(1);
2067    }
2068    dims.iter()
2069        .try_fold(1usize, |acc, &dim| acc.checked_mul(dim))
2070        .ok_or(StridedError::OffsetOverflow)
2071}
2072
2073fn checked_strided_offset(base: isize, strides: &[isize], index: &[usize]) -> Result<isize> {
2074    let mut offset = base;
2075    for (&stride, &coord) in strides.iter().zip(index.iter()) {
2076        offset = checked_offset_add(offset, stride, coord)?;
2077    }
2078    Ok(offset)
2079}
2080
2081fn checked_offset_add(base: isize, stride: isize, coord: usize) -> Result<isize> {
2082    let coord = isize::try_from(coord).map_err(|_| StridedError::OffsetOverflow)?;
2083    let scaled = stride
2084        .checked_mul(coord)
2085        .ok_or(StridedError::OffsetOverflow)?;
2086    base.checked_add(scaled).ok_or(StridedError::OffsetOverflow)
2087}
2088
2089struct CoordScratch {
2090    inline: [usize; RAW_FUSED_RANK_LIMIT],
2091    heap: Option<Vec<usize>>,
2092    len: usize,
2093}
2094
2095impl CoordScratch {
2096    fn new(len: usize) -> Self {
2097        if len <= RAW_FUSED_RANK_LIMIT {
2098            Self {
2099                inline: [0; RAW_FUSED_RANK_LIMIT],
2100                heap: None,
2101                len,
2102            }
2103        } else {
2104            Self {
2105                inline: [0; RAW_FUSED_RANK_LIMIT],
2106                heap: Some(vec![0; len]),
2107                len,
2108            }
2109        }
2110    }
2111
2112    fn as_mut_slice(&mut self) -> &mut [usize] {
2113        match &mut self.heap {
2114            Some(heap) => heap,
2115            None => &mut self.inline[..self.len],
2116        }
2117    }
2118}
2119
2120#[cfg(test)]
2121#[path = "gather_plan/tests/tests.rs"]
2122mod tests;