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