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