1use 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#[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
41pub 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#[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#[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#[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 window: Option<FusedPairLayout>,
393 #[cfg(not(feature = "parallel"))]
394 replay: WindowReplay,
395}
396
397#[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 window: Option<FusedPairLayout>,
418 #[cfg(not(feature = "parallel"))]
419 replay: WindowReplay,
420}
421
422#[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 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 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 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 for _ in 0..self.total {
707 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 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 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 for _ in 0..self.total {
771 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 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 for position in range {
814 let position = position as isize;
815 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 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 for _ in range {
872 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 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 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 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 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 for _ in 0..self.total {
1124 let value = unsafe { *operand_data.as_ptr().offset(state.source_offset) };
1126 unsafe { dest.write_at(state.dest_offset, value) };
1129 replay.advance(&mut state);
1130 }
1131 Ok(())
1132 }
1133
1134 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 let dest_ptr = unsafe { dest.data_ptr() }.cast::<MaybeUninit<T>>();
1164 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 let dest_ptr = unsafe { dest.data_ptr() };
1224 unsafe {
1227 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 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 for _ in range {
1278 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 #[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 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 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 for _ in 0..self.total {
1472 let value = unsafe { *update_data.as_ptr().offset(state.source_offset) };
1474 unsafe { dest.write_at(state.dest_offset, value) };
1477 replay.advance(&mut state);
1478 }
1479 Ok(())
1480 }
1481
1482 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 let dest_ptr = unsafe { dest.data_ptr() }.cast::<MaybeUninit<T>>();
1507 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 let dest_ptr = unsafe { dest.data_ptr() };
1567 unsafe {
1569 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 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 for _ in range {
1619 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 #[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 #[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 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 #[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 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 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 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 let dest_ptr = unsafe { dest.data_ptr() };
2008
2009 for _ in 0..self.batch_elems {
2012 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;