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