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 copy_plans: Vec<CopyPlan>,
102}
103
104impl SlicePlan {
105 #[allow(clippy::too_many_arguments)]
107 pub fn compile(
108 operand_dims: &[usize],
109 operand_strides: &[isize],
110 dest_dims: &[usize],
111 dest_strides: &[isize],
112 starts: &[usize],
113 limits: &[usize],
114 slice_strides: &[usize],
115 ) -> Result<Self> {
116 let rank = operand_dims.len();
117 if operand_strides.len() != rank || dest_dims.len() != rank || dest_strides.len() != rank {
118 return Err(StridedError::StrideLengthMismatch);
119 }
120 if starts.len() != rank {
121 return Err(StridedError::RankMismatch(starts.len(), rank));
122 }
123 if limits.len() != rank {
124 return Err(StridedError::RankMismatch(limits.len(), rank));
125 }
126 if slice_strides.len() != rank {
127 return Err(StridedError::RankMismatch(slice_strides.len(), rank));
128 }
129 checked_total_len(operand_dims)?;
130 checked_total_len(dest_dims)?;
131
132 let mut expected_dest_dims: AxisVec<usize> = AxisVec::with_capacity(rank);
133 let mut source_strides: AxisVec<isize> = AxisVec::with_capacity(rank);
134 let mut source_offset_delta = 0isize;
135 for axis in 0..rank {
136 let start = starts[axis];
137 let limit = limits[axis];
138 let stride = slice_strides[axis];
139 if start > limit || limit > operand_dims[axis] || stride == 0 {
140 return Err(StridedError::InvalidAxis { axis, rank });
141 }
142 let span = limit - start;
143 expected_dest_dims.push(span.div_ceil(stride));
144 source_strides.push(checked_stride_mul(operand_strides[axis], stride)?);
145 source_offset_delta =
146 checked_offset_add(source_offset_delta, operand_strides[axis], start)?;
147 }
148 if dest_dims != &expected_dest_dims[..] {
149 return Err(StridedError::ShapeMismatch(
150 dest_dims.to_vec(),
151 expected_dest_dims.to_vec(),
152 ));
153 }
154 let copy_plan = CopyPlan::compile(dest_dims, dest_strides, &source_strides)?;
155
156 Ok(Self {
157 operand_dims: operand_dims.into(),
158 operand_strides: operand_strides.into(),
159 dest_dims: dest_dims.into(),
160 dest_strides: dest_strides.into(),
161 source_strides,
162 source_offset_delta,
163 copy_plan,
164 })
165 }
166
167 pub fn execute<T>(
169 &self,
170 dest: &mut RawStridedMut<'_, T>,
171 operand: &RawStridedRef<'_, T>,
172 ) -> Result<()>
173 where
174 T: Copy + MaybeSendSync,
175 {
176 self.check_call(dest, operand)?;
177 let source_offset = operand
178 .offset()
179 .checked_add(self.source_offset_delta)
180 .ok_or(StridedError::OffsetOverflow)?;
181 let source = unsafe {
182 RawStridedRef::new_unchecked(
183 operand.data(),
184 &self.dest_dims,
185 &self.source_strides,
186 source_offset,
187 )
188 };
189 self.copy_plan.execute(dest, &source)
190 }
191
192 pub fn execute_uninit<T>(
194 &self,
195 dest: &mut RawStridedMut<'_, MaybeUninit<T>>,
196 operand: &RawStridedRef<'_, T>,
197 ) -> Result<()>
198 where
199 T: Copy + MaybeSendSync,
200 {
201 self.check_call(dest, operand)?;
202 let source_offset = operand
203 .offset()
204 .checked_add(self.source_offset_delta)
205 .ok_or(StridedError::OffsetOverflow)?;
206 let source = unsafe {
207 RawStridedRef::new_unchecked(
208 operand.data(),
209 &self.dest_dims,
210 &self.source_strides,
211 source_offset,
212 )
213 };
214 self.copy_plan.execute_uninit(dest, &source)
215 }
216
217 fn check_call<D, T>(
218 &self,
219 dest: &RawStridedMut<'_, D>,
220 operand: &RawStridedRef<'_, T>,
221 ) -> Result<()> {
222 if operand.dims() != &self.operand_dims[..]
223 || operand.strides() != &self.operand_strides[..]
224 || dest.dims() != &self.dest_dims[..]
225 || dest.strides() != &self.dest_strides[..]
226 {
227 return Err(StridedError::PlanLayoutMismatch);
228 }
229 Ok(())
230 }
231}
232
233impl PadPlan {
234 #[allow(clippy::too_many_arguments)]
236 pub fn compile(
237 operand_dims: &[usize],
238 operand_strides: &[isize],
239 dest_dims: &[usize],
240 dest_strides: &[isize],
241 edge_padding_low: &[i64],
242 edge_padding_high: &[i64],
243 interior_padding: &[i64],
244 ) -> Result<Self> {
245 let rank = operand_dims.len();
246 if operand_strides.len() != rank || dest_dims.len() != rank || dest_strides.len() != rank {
247 return Err(StridedError::StrideLengthMismatch);
248 }
249 if edge_padding_low.len() != rank {
250 return Err(StridedError::RankMismatch(edge_padding_low.len(), rank));
251 }
252 if edge_padding_high.len() != rank {
253 return Err(StridedError::RankMismatch(edge_padding_high.len(), rank));
254 }
255 if interior_padding.len() != rank {
256 return Err(StridedError::RankMismatch(interior_padding.len(), rank));
257 }
258
259 let operand_total = checked_total_len(operand_dims)?;
260 let dest_total = checked_total_len(dest_dims)?;
261 if !crate::layout_check::is_injective_layout(dest_dims, dest_strides) {
262 return Err(StridedError::NonInjectiveOutputLayout);
263 }
264
265 let mut expected_dest_dims: AxisVec<usize> = AxisVec::with_capacity(rank);
266 let mut interior_step: AxisVec<i64> = AxisVec::with_capacity(rank);
267 for axis in 0..rank {
268 if interior_padding[axis] < 0 {
269 return Err(StridedError::InvalidAxis { axis, rank });
270 }
271 let step = interior_padding[axis]
272 .checked_add(1)
273 .ok_or(StridedError::OffsetOverflow)?;
274 interior_step.push(step);
275 expected_dest_dims.push(checked_pad_output_dim(
276 operand_dims[axis],
277 edge_padding_low[axis],
278 edge_padding_high[axis],
279 step,
280 axis,
281 rank,
282 )?);
283 }
284 if dest_dims != &expected_dest_dims[..] {
285 return Err(StridedError::ShapeMismatch(
286 dest_dims.to_vec(),
287 expected_dest_dims.to_vec(),
288 ));
289 }
290 let contiguous_dest_fill = is_dense_col_major(dest_dims, dest_strides);
291 let contiguous_axis0_run = compile_contiguous_pad_axis0_run(
292 operand_dims,
293 operand_strides,
294 dest_dims,
295 dest_strides,
296 edge_padding_low,
297 &interior_step,
298 );
299 let generic_fill = compile_pad_fill_cursor(dest_dims, dest_strides)?;
300 let generic_copy = compile_pad_copy_cursor(
301 operand_dims,
302 operand_strides,
303 dest_dims,
304 dest_strides,
305 edge_padding_low,
306 &interior_step,
307 )?;
308
309 Ok(Self {
310 operand_dims: operand_dims.into(),
311 operand_strides: operand_strides.into(),
312 dest_dims: dest_dims.into(),
313 dest_strides: dest_strides.into(),
314 edge_padding_low: edge_padding_low.into(),
315 interior_step,
316 operand_total,
317 dest_total,
318 contiguous_dest_fill,
319 contiguous_axis0_run,
320 generic_fill,
321 generic_copy,
322 })
323 }
324
325 pub fn execute<T>(
327 &self,
328 dest: &mut RawStridedMut<'_, T>,
329 operand: &RawStridedRef<'_, T>,
330 fill: T,
331 ) -> Result<()>
332 where
333 T: Copy + MaybeSendSync,
334 {
335 self.check_call(dest, operand)?;
336 self.fill_dest(dest, fill)?;
337
338 if self.operand_total == 0 {
339 return Ok(());
340 }
341 if let Some(run) = self.contiguous_axis0_run {
342 return self.copy_operand_axis0_runs(dest, operand, run);
343 }
344 self.copy_operand(dest, operand)
345 }
346
347 pub fn execute_uninit<T>(
349 &self,
350 dest: &mut RawStridedMut<'_, MaybeUninit<T>>,
351 operand: &RawStridedRef<'_, T>,
352 fill: T,
353 ) -> Result<()>
354 where
355 T: Copy + MaybeSendSync,
356 {
357 self.check_call(dest, operand)?;
358 self.fill_dest(dest, MaybeUninit::new(fill))?;
359
360 if self.operand_total == 0 {
361 return Ok(());
362 }
363 if let Some(run) = self.contiguous_axis0_run {
364 return self.copy_operand_axis0_runs_uninit(dest, operand, run);
365 }
366 self.copy_operand_uninit(dest, operand)
367 }
368
369 fn fill_dest<T>(&self, dest: &mut RawStridedMut<'_, T>, fill: T) -> Result<()>
370 where
371 T: Copy + MaybeSendSync,
372 {
373 if self.dest_total == 0 {
374 return Ok(());
375 }
376 if self.contiguous_dest_fill {
377 let dest_offset =
378 usize::try_from(dest.offset()).map_err(|_| StridedError::OffsetOverflow)?;
379 let dest_end = dest_offset
380 .checked_add(self.dest_total)
381 .ok_or(StridedError::OffsetOverflow)?;
382 let dest_data = dest.data_mut();
383 let dest_slice = dest_data
384 .get_mut(dest_offset..dest_end)
385 .ok_or(StridedError::OffsetOverflow)?;
386 dest_slice.fill(fill);
387 return Ok(());
388 }
389 #[cfg(feature = "parallel")]
390 {
391 let nthreads = crate::threading::parallel_threads_for_len(self.dest_total);
392 if nthreads > 1 {
393 return self.fill_dest_parallel(dest, fill, nthreads);
394 }
395 }
396 self.fill_dest_serial(dest, fill)
397 }
398
399 fn copy_operand_axis0_runs<T>(
400 &self,
401 dest: &mut RawStridedMut<'_, T>,
402 operand: &RawStridedRef<'_, T>,
403 run: ContiguousPadAxis0Run,
404 ) -> Result<()>
405 where
406 T: Copy,
407 {
408 if run.len == 0 {
409 return Ok(());
410 }
411 let outer_dims = &self.operand_dims[1..];
412 let outer_total = checked_total_len(outer_dims)?;
413 let mut outer_idx_storage = CoordScratch::new(outer_dims.len());
414 let outer_idx = outer_idx_storage.as_mut_slice();
415 let operand_ptr = operand.data().as_ptr();
416 let dest_ptr = dest.data_mut().as_mut_ptr();
417
418 for _ in 0..outer_total {
419 let mut operand_offset =
420 checked_offset_add(operand.offset(), self.operand_strides[0], run.operand_start)?;
421 let mut dest_offset =
422 checked_offset_add(dest.offset(), self.dest_strides[0], run.dest_start)?;
423 let mut in_bounds = true;
424 for (outer_axis, &coord) in outer_idx.iter().enumerate() {
425 let axis = outer_axis + 1;
426 let out_pos = i128::from(self.edge_padding_low[axis])
427 + coord as i128 * i128::from(self.interior_step[axis]);
428 if out_pos < 0 || out_pos >= self.dest_dims[axis] as i128 {
429 in_bounds = false;
430 break;
431 }
432 operand_offset =
433 checked_offset_add(operand_offset, self.operand_strides[axis], coord)?;
434 dest_offset =
435 checked_offset_add(dest_offset, self.dest_strides[axis], out_pos as usize)?;
436 }
437 if in_bounds {
438 unsafe {
439 core::ptr::copy_nonoverlapping(
444 operand_ptr.offset(operand_offset),
445 dest_ptr.offset(dest_offset),
446 run.len,
447 );
448 }
449 }
450 advance_col_major_index(outer_idx, outer_dims);
451 }
452 Ok(())
453 }
454
455 fn copy_operand_axis0_runs_uninit<T>(
456 &self,
457 dest: &mut RawStridedMut<'_, MaybeUninit<T>>,
458 operand: &RawStridedRef<'_, T>,
459 run: ContiguousPadAxis0Run,
460 ) -> Result<()>
461 where
462 T: Copy,
463 {
464 if run.len == 0 {
465 return Ok(());
466 }
467 let outer_dims = &self.operand_dims[1..];
468 let outer_total = checked_total_len(outer_dims)?;
469 let mut outer_idx_storage = CoordScratch::new(outer_dims.len());
470 let outer_idx = outer_idx_storage.as_mut_slice();
471 let operand_ptr = operand.data().as_ptr();
472 let dest_ptr = dest.data_mut().as_mut_ptr();
473
474 for _ in 0..outer_total {
475 let mut operand_offset =
476 checked_offset_add(operand.offset(), self.operand_strides[0], run.operand_start)?;
477 let mut dest_offset =
478 checked_offset_add(dest.offset(), self.dest_strides[0], run.dest_start)?;
479 let mut in_bounds = true;
480 for (outer_axis, &coord) in outer_idx.iter().enumerate() {
481 let axis = outer_axis + 1;
482 let out_pos = i128::from(self.edge_padding_low[axis])
483 + coord as i128 * i128::from(self.interior_step[axis]);
484 if out_pos < 0 || out_pos >= self.dest_dims[axis] as i128 {
485 in_bounds = false;
486 break;
487 }
488 operand_offset =
489 checked_offset_add(operand_offset, self.operand_strides[axis], coord)?;
490 dest_offset =
491 checked_offset_add(dest_offset, self.dest_strides[axis], out_pos as usize)?;
492 }
493 if in_bounds {
494 unsafe {
495 core::ptr::copy_nonoverlapping(
496 operand_ptr.offset(operand_offset),
497 dest_ptr.offset(dest_offset).cast::<T>(),
498 run.len,
499 );
500 }
501 }
502 advance_col_major_index(outer_idx, outer_dims);
503 }
504 Ok(())
505 }
506
507 #[cfg(test)]
508 fn contiguous_axis0_run(&self) -> Option<(usize, usize, usize)> {
509 self.contiguous_axis0_run
510 .map(|run| (run.operand_start, run.dest_start, run.len))
511 }
512
513 #[cfg(test)]
514 fn has_contiguous_dest_fill(&self) -> bool {
515 self.contiguous_dest_fill
516 }
517
518 fn fill_dest_serial<T>(&self, dest: &mut RawStridedMut<'_, T>, fill: T) -> Result<()>
519 where
520 T: Copy,
521 {
522 let dest_ptr = dest.data_mut().as_mut_ptr();
523 let mut cursor = PadFillState::new(dest.offset(), &self.dest_dims, &self.generic_fill);
524 for _ in 0..self.dest_total {
525 unsafe {
526 *dest_ptr.offset(cursor.offset) = fill;
530 }
531 cursor.advance(&self.dest_dims, &self.generic_fill);
532 }
533 Ok(())
534 }
535
536 #[cfg(feature = "parallel")]
537 fn fill_dest_parallel<T>(
538 &self,
539 dest: &mut RawStridedMut<'_, T>,
540 fill: T,
541 nthreads: usize,
542 ) -> Result<()>
543 where
544 T: Copy + MaybeSendSync,
545 {
546 let dest_offset_base = dest.offset();
547 let dest_ptr = crate::threading::SendPtr(dest.data_mut().as_mut_ptr());
548 crate::threading::parallel_map_reduce(
549 0..self.dest_total,
550 nthreads,
551 &|range| {
552 let mut cursor = PadFillState::decode(
553 range.start,
554 dest_offset_base,
555 &self.dest_dims,
556 &self.generic_fill,
557 )?;
558 let dest_ptr = dest_ptr.as_ptr();
559 for _ in range {
560 unsafe {
561 *dest_ptr.offset(cursor.offset) = fill;
566 }
567 cursor.advance(&self.dest_dims, &self.generic_fill);
568 }
569 Ok(())
570 },
571 &|left, right| left.and(right),
572 )
573 }
574
575 fn copy_operand<T>(
576 &self,
577 dest: &mut RawStridedMut<'_, T>,
578 operand: &RawStridedRef<'_, T>,
579 ) -> Result<()>
580 where
581 T: Copy + MaybeSendSync,
582 {
583 #[cfg(feature = "parallel")]
584 {
585 let nthreads = crate::threading::parallel_threads_for_len(self.generic_copy.total);
586 if nthreads > 1 {
587 return self.copy_operand_parallel(dest, operand, nthreads);
588 }
589 }
590 self.copy_operand_serial(dest, operand)
591 }
592
593 fn copy_operand_uninit<T>(
594 &self,
595 dest: &mut RawStridedMut<'_, MaybeUninit<T>>,
596 operand: &RawStridedRef<'_, T>,
597 ) -> Result<()>
598 where
599 T: Copy + MaybeSendSync,
600 {
601 #[cfg(feature = "parallel")]
602 {
603 let nthreads = crate::threading::parallel_threads_for_len(self.generic_copy.total);
604 if nthreads > 1 {
605 return self.copy_operand_uninit_parallel(dest, operand, nthreads);
606 }
607 }
608 self.copy_operand_uninit_serial(dest, operand)
609 }
610
611 fn copy_operand_uninit_serial<T>(
612 &self,
613 dest: &mut RawStridedMut<'_, MaybeUninit<T>>,
614 operand: &RawStridedRef<'_, T>,
615 ) -> Result<()>
616 where
617 T: Copy,
618 {
619 if self.generic_copy.total == 0 {
620 return Ok(());
621 }
622 let operand_ptr = operand.data().as_ptr();
623 let dest_ptr = dest.data_mut().as_mut_ptr();
624 let source_base = operand
625 .offset()
626 .checked_add(self.generic_copy.source_base_delta)
627 .ok_or(StridedError::OffsetOverflow)?;
628 let dest_base = dest
629 .offset()
630 .checked_add(self.generic_copy.dest_base_delta)
631 .ok_or(StridedError::OffsetOverflow)?;
632 let mut cursor = PadCopyState::new(source_base, dest_base, &self.generic_copy);
633 for _ in 0..self.generic_copy.total {
634 unsafe {
635 (*dest_ptr.offset(cursor.dest_offset))
640 .write(*operand_ptr.offset(cursor.source_offset));
641 }
642 cursor.advance(&self.generic_copy);
643 }
644 Ok(())
645 }
646
647 #[cfg(feature = "parallel")]
648 fn copy_operand_uninit_parallel<T>(
649 &self,
650 dest: &mut RawStridedMut<'_, MaybeUninit<T>>,
651 operand: &RawStridedRef<'_, T>,
652 nthreads: usize,
653 ) -> Result<()>
654 where
655 T: Copy + MaybeSendSync,
656 {
657 let operand_base = operand
658 .offset()
659 .checked_add(self.generic_copy.source_base_delta)
660 .ok_or(StridedError::OffsetOverflow)?;
661 let operand_ptr = crate::threading::SendPtr(operand.data().as_ptr() as *mut T);
662 let dest_base = dest
663 .offset()
664 .checked_add(self.generic_copy.dest_base_delta)
665 .ok_or(StridedError::OffsetOverflow)?;
666 let dest_ptr = crate::threading::SendPtr(dest.data_mut().as_mut_ptr());
667 crate::threading::parallel_map_reduce(
668 0..self.generic_copy.total,
669 nthreads,
670 &|range| {
671 let mut cursor =
672 PadCopyState::decode(range.start, operand_base, dest_base, &self.generic_copy)?;
673 let operand_ptr = operand_ptr.as_const();
674 let dest_ptr = dest_ptr.as_ptr();
675
676 for _ in range {
677 unsafe {
678 (*dest_ptr.offset(cursor.dest_offset))
683 .write(*operand_ptr.offset(cursor.source_offset));
684 }
685 cursor.advance(&self.generic_copy);
686 }
687 Ok(())
688 },
689 &|left, right| left.and(right),
690 )
691 }
692
693 fn copy_operand_serial<T>(
694 &self,
695 dest: &mut RawStridedMut<'_, T>,
696 operand: &RawStridedRef<'_, T>,
697 ) -> Result<()>
698 where
699 T: Copy,
700 {
701 if self.generic_copy.total == 0 {
702 return Ok(());
703 }
704 let operand_ptr = operand.data().as_ptr();
705 let dest_ptr = dest.data_mut().as_mut_ptr();
706 let source_base = operand
707 .offset()
708 .checked_add(self.generic_copy.source_base_delta)
709 .ok_or(StridedError::OffsetOverflow)?;
710 let dest_base = dest
711 .offset()
712 .checked_add(self.generic_copy.dest_base_delta)
713 .ok_or(StridedError::OffsetOverflow)?;
714 let mut cursor = PadCopyState::new(source_base, dest_base, &self.generic_copy);
715 for _ in 0..self.generic_copy.total {
716 unsafe {
717 *dest_ptr.offset(cursor.dest_offset) = *operand_ptr.offset(cursor.source_offset);
722 }
723 cursor.advance(&self.generic_copy);
724 }
725 Ok(())
726 }
727
728 #[cfg(feature = "parallel")]
729 fn copy_operand_parallel<T>(
730 &self,
731 dest: &mut RawStridedMut<'_, T>,
732 operand: &RawStridedRef<'_, T>,
733 nthreads: usize,
734 ) -> Result<()>
735 where
736 T: Copy + MaybeSendSync,
737 {
738 let operand_base = operand
739 .offset()
740 .checked_add(self.generic_copy.source_base_delta)
741 .ok_or(StridedError::OffsetOverflow)?;
742 let operand_ptr = crate::threading::SendPtr(operand.data().as_ptr() as *mut T);
743 let dest_base = dest
744 .offset()
745 .checked_add(self.generic_copy.dest_base_delta)
746 .ok_or(StridedError::OffsetOverflow)?;
747 let dest_ptr = crate::threading::SendPtr(dest.data_mut().as_mut_ptr());
748 crate::threading::parallel_map_reduce(
749 0..self.generic_copy.total,
750 nthreads,
751 &|range| {
752 let mut cursor =
753 PadCopyState::decode(range.start, operand_base, dest_base, &self.generic_copy)?;
754 let operand_ptr = operand_ptr.as_const();
755 let dest_ptr = dest_ptr.as_ptr();
756
757 for _ in range {
758 unsafe {
759 *dest_ptr.offset(cursor.dest_offset) =
764 *operand_ptr.offset(cursor.source_offset);
765 }
766 cursor.advance(&self.generic_copy);
767 }
768 Ok(())
769 },
770 &|left, right| left.and(right),
771 )
772 }
773
774 fn check_call<D, T>(
775 &self,
776 dest: &RawStridedMut<'_, D>,
777 operand: &RawStridedRef<'_, T>,
778 ) -> Result<()> {
779 if operand.dims() != &self.operand_dims[..]
780 || operand.strides() != &self.operand_strides[..]
781 || dest.dims() != &self.dest_dims[..]
782 || dest.strides() != &self.dest_strides[..]
783 {
784 return Err(StridedError::PlanLayoutMismatch);
785 }
786 Ok(())
787 }
788}
789
790impl ConcatenatePlan {
791 pub fn compile(
793 input_dims: &[&[usize]],
794 input_strides: &[&[isize]],
795 dest_dims: &[usize],
796 dest_strides: &[isize],
797 axis: usize,
798 ) -> Result<Self> {
799 if input_dims.is_empty() {
800 return Err(StridedError::UnsupportedArity {
801 arity: 0,
802 max: usize::MAX,
803 });
804 }
805 if input_dims.len() != input_strides.len() {
806 return Err(StridedError::RankMismatch(
807 input_strides.len(),
808 input_dims.len(),
809 ));
810 }
811
812 let rank = input_dims[0].len();
813 if dest_dims.len() != rank || dest_strides.len() != rank {
814 return Err(StridedError::StrideLengthMismatch);
815 }
816 if axis >= rank {
817 return Err(StridedError::InvalidAxis { axis, rank });
818 }
819 checked_total_len(dest_dims)?;
820 if !crate::layout_check::is_injective_layout(dest_dims, dest_strides) {
821 return Err(StridedError::NonInjectiveOutputLayout);
822 }
823
824 let mut expected_dest_dims: AxisVec<usize> = input_dims[0].into();
825 expected_dest_dims[axis] = 0;
826 let mut stored_input_dims = Vec::with_capacity(input_dims.len());
827 let mut stored_input_strides = Vec::with_capacity(input_dims.len());
828 let mut dest_offset_deltas = Vec::with_capacity(input_dims.len());
829 let mut copy_plans = Vec::with_capacity(input_dims.len());
830 let mut axis_base = 0usize;
831
832 for (dims, strides) in input_dims.iter().zip(input_strides.iter()) {
833 if dims.len() != rank {
834 return Err(StridedError::RankMismatch(dims.len(), rank));
835 }
836 if strides.len() != rank {
837 return Err(StridedError::StrideLengthMismatch);
838 }
839 checked_total_len(dims)?;
840 for dim in 0..rank {
841 if dim == axis {
842 expected_dest_dims[axis] = expected_dest_dims[axis]
843 .checked_add(dims[axis])
844 .ok_or(StridedError::OffsetOverflow)?;
845 } else if dims[dim] != input_dims[0][dim] {
846 return Err(StridedError::ShapeMismatch(
847 dims.to_vec(),
848 input_dims[0].to_vec(),
849 ));
850 }
851 }
852 dest_offset_deltas.push(checked_offset_add(0, dest_strides[axis], axis_base)?);
853 axis_base = axis_base
854 .checked_add(dims[axis])
855 .ok_or(StridedError::OffsetOverflow)?;
856 copy_plans.push(CopyPlan::compile(dims, dest_strides, strides)?);
857 stored_input_dims.push((*dims).into());
858 stored_input_strides.push((*strides).into());
859 }
860
861 if dest_dims != &expected_dest_dims[..] {
862 return Err(StridedError::ShapeMismatch(
863 dest_dims.to_vec(),
864 expected_dest_dims.to_vec(),
865 ));
866 }
867
868 Ok(Self {
869 input_dims: stored_input_dims,
870 input_strides: stored_input_strides,
871 dest_dims: dest_dims.into(),
872 dest_strides: dest_strides.into(),
873 dest_offset_deltas,
874 copy_plans,
875 })
876 }
877
878 pub fn execute<T>(
880 &self,
881 dest: &mut RawStridedMut<'_, T>,
882 inputs: &[RawStridedRef<'_, T>],
883 ) -> Result<()>
884 where
885 T: Copy + MaybeSendSync,
886 {
887 self.check_dest_layout(dest)?;
888 if inputs.len() != self.input_dims.len() {
889 return Err(StridedError::RankMismatch(
890 inputs.len(),
891 self.input_dims.len(),
892 ));
893 }
894 for (position, input) in inputs.iter().enumerate() {
895 self.check_input_layout(position, input)?;
896 self.segment_offset(position, dest.offset())?;
897 }
898 for (position, input) in inputs.iter().enumerate() {
899 self.execute_segment(position, dest, input)?;
900 }
901 Ok(())
902 }
903
904 pub fn execute_uninit<T>(
906 &self,
907 dest: &mut RawStridedMut<'_, MaybeUninit<T>>,
908 inputs: &[RawStridedRef<'_, T>],
909 ) -> Result<()>
910 where
911 T: Copy + MaybeSendSync,
912 {
913 self.check_dest_layout(dest)?;
914 if inputs.len() != self.input_dims.len() {
915 return Err(StridedError::RankMismatch(
916 inputs.len(),
917 self.input_dims.len(),
918 ));
919 }
920 for (position, input) in inputs.iter().enumerate() {
921 self.check_input_layout(position, input)?;
922 self.segment_offset(position, dest.offset())?;
923 }
924 for (position, input) in inputs.iter().enumerate() {
925 self.execute_segment_uninit(position, dest, input)?;
926 }
927 Ok(())
928 }
929
930 pub(crate) fn check_dest_layout<T>(&self, dest: &RawStridedMut<'_, T>) -> Result<()> {
931 if dest.dims() != &self.dest_dims[..] || dest.strides() != &self.dest_strides[..] {
932 return Err(StridedError::PlanLayoutMismatch);
933 }
934 Ok(())
935 }
936
937 pub(crate) fn check_input_layout<T>(
938 &self,
939 position: usize,
940 input: &RawStridedRef<'_, T>,
941 ) -> Result<()> {
942 if position >= self.input_dims.len()
943 || input.dims() != &self.input_dims[position][..]
944 || input.strides() != &self.input_strides[position][..]
945 {
946 return Err(StridedError::PlanLayoutMismatch);
947 }
948 Ok(())
949 }
950
951 pub(crate) fn input_count(&self) -> usize {
952 self.input_dims.len()
953 }
954
955 pub(crate) fn segment_offset(&self, position: usize, dest_offset: isize) -> Result<isize> {
956 dest_offset
957 .checked_add(self.dest_offset_deltas[position])
958 .ok_or(StridedError::OffsetOverflow)
959 }
960
961 pub(crate) fn execute_segment<T>(
962 &self,
963 position: usize,
964 dest: &mut RawStridedMut<'_, T>,
965 input: &RawStridedRef<'_, T>,
966 ) -> Result<()>
967 where
968 T: Copy + MaybeSendSync,
969 {
970 let segment_offset = self.segment_offset(position, dest.offset())?;
971 let dest_data = dest.data_mut();
972 let mut segment = unsafe {
973 RawStridedMut::new_unchecked(
974 dest_data,
975 &self.input_dims[position],
976 &self.dest_strides,
977 segment_offset,
978 )
979 };
980 self.copy_plans[position].execute(&mut segment, input)
981 }
982
983 pub(crate) fn execute_segment_uninit<T>(
984 &self,
985 position: usize,
986 dest: &mut RawStridedMut<'_, MaybeUninit<T>>,
987 input: &RawStridedRef<'_, T>,
988 ) -> Result<()>
989 where
990 T: Copy + MaybeSendSync,
991 {
992 let segment_offset = self.segment_offset(position, dest.offset())?;
993 let dest_data = dest.data_mut();
994 let mut segment = unsafe {
995 RawStridedMut::new_unchecked(
996 dest_data,
997 &self.input_dims[position],
998 &self.dest_strides,
999 segment_offset,
1000 )
1001 };
1002 self.copy_plans[position].execute_uninit(&mut segment, input)
1003 }
1004}
1005
1006impl ReversePlan {
1007 pub fn compile(
1009 operand_dims: &[usize],
1010 operand_strides: &[isize],
1011 dest_strides: &[isize],
1012 axes: &[usize],
1013 ) -> Result<Self> {
1014 let rank = operand_dims.len();
1015 if operand_strides.len() != rank || dest_strides.len() != rank {
1016 return Err(StridedError::StrideLengthMismatch);
1017 }
1018 checked_total_len(operand_dims)?;
1019
1020 let mut reverse_axis: AxisVec<bool> = (0..rank).map(|_| false).collect();
1021 for &axis in axes {
1022 if axis >= rank {
1023 return Err(StridedError::InvalidAxis { axis, rank });
1024 }
1025 reverse_axis[axis] = true;
1026 }
1027
1028 let mut source_strides: AxisVec<isize> = AxisVec::with_capacity(rank);
1029 let mut source_offset_delta = 0isize;
1030 for axis in 0..rank {
1031 if reverse_axis[axis] {
1032 source_strides.push(
1033 operand_strides[axis]
1034 .checked_neg()
1035 .ok_or(StridedError::OffsetOverflow)?,
1036 );
1037 if operand_dims[axis] > 0 {
1038 source_offset_delta = checked_offset_add(
1039 source_offset_delta,
1040 operand_strides[axis],
1041 operand_dims[axis] - 1,
1042 )?;
1043 }
1044 } else {
1045 source_strides.push(operand_strides[axis]);
1046 }
1047 }
1048 let copy_plan = CopyPlan::compile(operand_dims, dest_strides, &source_strides)?;
1049
1050 Ok(Self {
1051 operand_dims: operand_dims.into(),
1052 operand_strides: operand_strides.into(),
1053 dest_strides: dest_strides.into(),
1054 source_strides,
1055 source_offset_delta,
1056 copy_plan,
1057 })
1058 }
1059
1060 pub fn execute<T>(
1062 &self,
1063 dest: &mut RawStridedMut<'_, T>,
1064 operand: &RawStridedRef<'_, T>,
1065 ) -> Result<()>
1066 where
1067 T: Copy + MaybeSendSync,
1068 {
1069 self.check_call(dest, operand)?;
1070 let source_offset = operand
1071 .offset()
1072 .checked_add(self.source_offset_delta)
1073 .ok_or(StridedError::OffsetOverflow)?;
1074 let source = unsafe {
1075 RawStridedRef::new_unchecked(
1076 operand.data(),
1077 &self.operand_dims,
1078 &self.source_strides,
1079 source_offset,
1080 )
1081 };
1082 self.copy_plan.execute(dest, &source)
1083 }
1084
1085 pub fn execute_uninit<T>(
1087 &self,
1088 dest: &mut RawStridedMut<'_, MaybeUninit<T>>,
1089 operand: &RawStridedRef<'_, T>,
1090 ) -> Result<()>
1091 where
1092 T: Copy + MaybeSendSync,
1093 {
1094 self.check_call(dest, operand)?;
1095 let source_offset = operand
1096 .offset()
1097 .checked_add(self.source_offset_delta)
1098 .ok_or(StridedError::OffsetOverflow)?;
1099 let source = unsafe {
1100 RawStridedRef::new_unchecked(
1101 operand.data(),
1102 &self.operand_dims,
1103 &self.source_strides,
1104 source_offset,
1105 )
1106 };
1107 self.copy_plan.execute_uninit(dest, &source)
1108 }
1109
1110 fn check_call<D, T>(
1111 &self,
1112 dest: &RawStridedMut<'_, D>,
1113 operand: &RawStridedRef<'_, T>,
1114 ) -> Result<()> {
1115 if operand.dims() != &self.operand_dims[..]
1116 || operand.strides() != &self.operand_strides[..]
1117 || dest.dims() != &self.operand_dims[..]
1118 || dest.strides() != &self.dest_strides[..]
1119 {
1120 return Err(StridedError::PlanLayoutMismatch);
1121 }
1122 Ok(())
1123 }
1124}
1125
1126fn checked_total_len(dims: &[usize]) -> Result<usize> {
1127 if dims.is_empty() {
1128 return Ok(1);
1129 }
1130 dims.iter()
1131 .try_fold(1usize, |acc, &dim| acc.checked_mul(dim))
1132 .ok_or(StridedError::OffsetOverflow)
1133}
1134
1135fn checked_stride_mul(stride: isize, factor: usize) -> Result<isize> {
1136 let factor = isize::try_from(factor).map_err(|_| StridedError::OffsetOverflow)?;
1137 stride
1138 .checked_mul(factor)
1139 .ok_or(StridedError::OffsetOverflow)
1140}
1141
1142fn checked_pad_output_dim(
1143 input_extent: usize,
1144 edge_low: i64,
1145 edge_high: i64,
1146 interior_step: i64,
1147 axis: usize,
1148 rank: usize,
1149) -> Result<usize> {
1150 let base = if input_extent == 0 {
1151 0i128
1152 } else {
1153 (input_extent as i128 - 1)
1154 .checked_mul(i128::from(interior_step))
1155 .and_then(|value| value.checked_add(1))
1156 .ok_or(StridedError::OffsetOverflow)?
1157 };
1158 let dim = i128::from(edge_low)
1159 .checked_add(i128::from(edge_high))
1160 .and_then(|value| value.checked_add(base))
1161 .ok_or(StridedError::OffsetOverflow)?;
1162 usize::try_from(dim).map_err(|_| StridedError::InvalidAxis { axis, rank })
1163}
1164
1165fn compile_pad_fill_cursor(dims: &[usize], strides: &[isize]) -> Result<PadFillCursor> {
1166 let mut steps = AxisVec::with_capacity(dims.len());
1167 let mut resets = AxisVec::with_capacity(dims.len());
1168 for (&dim, &stride) in dims.iter().zip(strides) {
1169 steps.push(stride);
1170 resets.push(checked_cursor_reset(stride, dim)?);
1171 }
1172 check_offset_span(dims, &steps)?;
1173 Ok(PadFillCursor { steps, resets })
1174}
1175
1176fn compile_pad_copy_cursor(
1177 operand_dims: &[usize],
1178 operand_strides: &[isize],
1179 dest_dims: &[usize],
1180 dest_strides: &[isize],
1181 edge_padding_low: &[i64],
1182 interior_step: &[i64],
1183) -> Result<PadCopyCursor> {
1184 let mut shape = AxisVec::with_capacity(operand_dims.len());
1185 let mut source_steps = AxisVec::with_capacity(operand_dims.len());
1186 let mut source_resets = AxisVec::with_capacity(operand_dims.len());
1187 let mut dest_steps = AxisVec::with_capacity(operand_dims.len());
1188 let mut dest_resets = AxisVec::with_capacity(operand_dims.len());
1189 let mut source_base_delta = 0isize;
1190 let mut dest_base_delta = 0isize;
1191 let mut copy_empty = false;
1192
1193 for axis in 0..operand_dims.len() {
1194 let (start, end) = checked_pad_valid_interval(
1195 operand_dims[axis],
1196 dest_dims[axis],
1197 edge_padding_low[axis],
1198 interior_step[axis],
1199 )?;
1200 let extent = end - start;
1201 shape.push(extent);
1202 if !copy_empty && extent != 0 {
1203 source_base_delta =
1204 checked_offset_add(source_base_delta, operand_strides[axis], start)?;
1205 let output_start = i128::from(edge_padding_low[axis])
1206 .checked_add(
1207 i128::try_from(start)
1208 .map_err(|_| StridedError::OffsetOverflow)?
1209 .checked_mul(i128::from(interior_step[axis]))
1210 .ok_or(StridedError::OffsetOverflow)?,
1211 )
1212 .ok_or(StridedError::OffsetOverflow)?;
1213 let output_start =
1214 usize::try_from(output_start).map_err(|_| StridedError::OffsetOverflow)?;
1215 dest_base_delta =
1216 checked_offset_add(dest_base_delta, dest_strides[axis], output_start)?;
1217 }
1218 copy_empty |= extent == 0;
1219
1220 let source_step = operand_strides[axis];
1221 let dest_step = checked_stride_mul_i64(dest_strides[axis], interior_step[axis])?;
1222 source_steps.push(source_step);
1223 source_resets.push(checked_cursor_reset(source_step, extent)?);
1224 dest_steps.push(dest_step);
1225 dest_resets.push(checked_cursor_reset(dest_step, extent)?);
1226 }
1227
1228 check_offset_span(&shape, &source_steps)?;
1229 check_offset_span(&shape, &dest_steps)?;
1230 let total = if shape.iter().any(|&extent| extent == 0) {
1231 0
1232 } else {
1233 checked_total_len(&shape)?
1234 };
1235
1236 Ok(PadCopyCursor {
1237 shape,
1238 source_base_delta,
1239 dest_base_delta,
1240 source_steps,
1241 source_resets,
1242 dest_steps,
1243 dest_resets,
1244 total,
1245 })
1246}
1247
1248fn checked_pad_valid_interval(
1249 input_extent: usize,
1250 dest_extent: usize,
1251 edge_low: i64,
1252 step: i64,
1253) -> Result<(usize, usize)> {
1254 if input_extent == 0 || dest_extent == 0 {
1255 return Ok((0, 0));
1256 }
1257 let step = i128::from(step);
1258 let lower = ceil_div_positive(-i128::from(edge_low), step)?;
1259 let dest_last = i128::try_from(dest_extent)
1260 .map_err(|_| StridedError::OffsetOverflow)?
1261 .checked_sub(1)
1262 .ok_or(StridedError::OffsetOverflow)?;
1263 let upper = floor_div_positive(
1264 dest_last
1265 .checked_sub(i128::from(edge_low))
1266 .ok_or(StridedError::OffsetOverflow)?,
1267 step,
1268 )?
1269 .checked_add(1)
1270 .ok_or(StridedError::OffsetOverflow)?;
1271 let input_extent = i128::try_from(input_extent).map_err(|_| StridedError::OffsetOverflow)?;
1272 let lower = lower.clamp(0, input_extent);
1273 let upper = upper.clamp(0, input_extent);
1274 if lower >= upper {
1275 return Ok((0, 0));
1276 }
1277 Ok((
1278 usize::try_from(lower).map_err(|_| StridedError::OffsetOverflow)?,
1279 usize::try_from(upper).map_err(|_| StridedError::OffsetOverflow)?,
1280 ))
1281}
1282
1283fn ceil_div_positive(numerator: i128, denominator: i128) -> Result<i128> {
1284 if denominator <= 0 {
1285 return Err(StridedError::OffsetOverflow);
1286 }
1287 let quotient = numerator.div_euclid(denominator);
1288 let remainder = numerator.rem_euclid(denominator);
1289 quotient
1290 .checked_add(i128::from(remainder != 0))
1291 .ok_or(StridedError::OffsetOverflow)
1292}
1293
1294fn floor_div_positive(numerator: i128, denominator: i128) -> Result<i128> {
1295 if denominator <= 0 {
1296 return Err(StridedError::OffsetOverflow);
1297 }
1298 Ok(numerator.div_euclid(denominator))
1299}
1300
1301fn checked_stride_mul_i64(stride: isize, factor: i64) -> Result<isize> {
1302 let factor = isize::try_from(factor).map_err(|_| StridedError::OffsetOverflow)?;
1303 stride
1304 .checked_mul(factor)
1305 .ok_or(StridedError::OffsetOverflow)
1306}
1307
1308fn checked_cursor_reset(step: isize, extent: usize) -> Result<isize> {
1309 if extent == 0 {
1310 return Ok(0);
1311 }
1312 let last = isize::try_from(extent - 1).map_err(|_| StridedError::OffsetOverflow)?;
1313 step.checked_mul(last)
1314 .and_then(isize::checked_neg)
1315 .ok_or(StridedError::OffsetOverflow)
1316}
1317
1318fn check_offset_span(shape: &[usize], strides: &[isize]) -> Result<()> {
1319 let mut min = 0isize;
1320 let mut max = 0isize;
1321 for (&extent, &stride) in shape.iter().zip(strides) {
1322 if extent <= 1 {
1323 continue;
1324 }
1325 let last = isize::try_from(extent - 1).map_err(|_| StridedError::OffsetOverflow)?;
1326 let delta = stride
1327 .checked_mul(last)
1328 .ok_or(StridedError::OffsetOverflow)?;
1329 if delta < 0 {
1330 min = min.checked_add(delta).ok_or(StridedError::OffsetOverflow)?;
1331 } else {
1332 max = max.checked_add(delta).ok_or(StridedError::OffsetOverflow)?;
1333 }
1334 }
1335 let _ = (min, max);
1336 Ok(())
1337}
1338
1339fn is_dense_col_major(dims: &[usize], strides: &[isize]) -> bool {
1340 let mut expected = 1isize;
1341 for (&dim, &stride) in dims.iter().zip(strides.iter()) {
1342 if stride != expected {
1343 return false;
1344 }
1345 let Ok(dim) = isize::try_from(dim) else {
1346 return false;
1347 };
1348 let Some(next) = expected.checked_mul(dim) else {
1349 return false;
1350 };
1351 expected = next;
1352 }
1353 true
1354}
1355
1356fn compile_contiguous_pad_axis0_run(
1357 operand_dims: &[usize],
1358 operand_strides: &[isize],
1359 dest_dims: &[usize],
1360 dest_strides: &[isize],
1361 edge_padding_low: &[i64],
1362 interior_step: &[i64],
1363) -> Option<ContiguousPadAxis0Run> {
1364 if operand_dims.is_empty()
1365 || operand_strides[0] != 1
1366 || dest_strides[0] != 1
1367 || interior_step[0] != 1
1368 {
1369 return None;
1370 }
1371
1372 let operand_extent = operand_dims[0] as i128;
1373 let dest_extent = dest_dims[0] as i128;
1374 let edge_low = i128::from(edge_padding_low[0]);
1375 let operand_start = (-edge_low).clamp(0, operand_extent);
1376 let dest_start = edge_low.clamp(0, dest_extent);
1377 let len = (operand_extent - operand_start).min(dest_extent - dest_start);
1378 Some(ContiguousPadAxis0Run {
1379 operand_start: usize::try_from(operand_start).ok()?,
1380 dest_start: usize::try_from(dest_start).ok()?,
1381 len: usize::try_from(len).ok()?,
1382 })
1383}
1384
1385fn checked_offset_add(base: isize, stride: isize, coord: usize) -> Result<isize> {
1386 let coord = isize::try_from(coord).map_err(|_| StridedError::OffsetOverflow)?;
1387 let scaled = stride
1388 .checked_mul(coord)
1389 .ok_or(StridedError::OffsetOverflow)?;
1390 base.checked_add(scaled).ok_or(StridedError::OffsetOverflow)
1391}
1392
1393fn advance_col_major_index(index: &mut [usize], shape: &[usize]) {
1394 for axis in 0..index.len() {
1395 index[axis] += 1;
1396 if index[axis] < shape[axis] {
1397 return;
1398 }
1399 index[axis] = 0;
1400 }
1401}
1402
1403#[cfg(feature = "parallel")]
1404fn fill_col_major_index(mut linear: usize, shape: &[usize], out: &mut [usize]) {
1405 for (axis, coord) in out.iter_mut().enumerate() {
1406 let dim = shape[axis];
1407 *coord = linear % dim;
1408 linear /= dim;
1409 }
1410}
1411
1412struct PadFillState {
1413 coords: CoordScratch,
1414 offset: isize,
1415}
1416
1417impl PadFillState {
1418 fn new(base: isize, shape: &[usize], _cursor: &PadFillCursor) -> Self {
1419 Self {
1420 coords: CoordScratch::new(shape.len()),
1421 offset: base,
1422 }
1423 }
1424
1425 #[cfg(feature = "parallel")]
1426 fn decode(linear: usize, base: isize, shape: &[usize], cursor: &PadFillCursor) -> Result<Self> {
1427 let mut state = Self::new(base, shape, cursor);
1428 fill_col_major_index(linear, shape, state.coords.as_mut_slice());
1429 for (&coord, &step) in state.coords.as_mut_slice().iter().zip(&cursor.steps) {
1430 state.offset = checked_offset_add(state.offset, step, coord)?;
1431 }
1432 Ok(state)
1433 }
1434
1435 #[inline]
1436 fn advance(&mut self, shape: &[usize], cursor: &PadFillCursor) {
1437 for axis in 0..shape.len() {
1441 let next = self.coords.as_mut_slice()[axis] + 1;
1442 if next < shape[axis] {
1443 self.coords.as_mut_slice()[axis] = next;
1444 self.offset += cursor.steps[axis];
1445 return;
1446 }
1447 self.coords.as_mut_slice()[axis] = 0;
1448 self.offset += cursor.resets[axis];
1449 }
1450 }
1451}
1452
1453struct PadCopyState {
1454 coords: CoordScratch,
1455 source_offset: isize,
1456 dest_offset: isize,
1457}
1458
1459impl PadCopyState {
1460 fn new(source_base: isize, dest_base: isize, cursor: &PadCopyCursor) -> Self {
1461 Self {
1462 coords: CoordScratch::new(cursor.shape.len()),
1463 source_offset: source_base,
1464 dest_offset: dest_base,
1465 }
1466 }
1467
1468 #[cfg(feature = "parallel")]
1469 fn decode(
1470 linear: usize,
1471 source_base: isize,
1472 dest_base: isize,
1473 cursor: &PadCopyCursor,
1474 ) -> Result<Self> {
1475 let mut state = Self::new(source_base, dest_base, cursor);
1476 fill_col_major_index(linear, &cursor.shape, state.coords.as_mut_slice());
1477 for axis in 0..cursor.shape.len() {
1478 let coord = state.coords.as_mut_slice()[axis];
1479 state.source_offset =
1480 checked_offset_add(state.source_offset, cursor.source_steps[axis], coord)?;
1481 state.dest_offset =
1482 checked_offset_add(state.dest_offset, cursor.dest_steps[axis], coord)?;
1483 }
1484 Ok(state)
1485 }
1486
1487 #[inline]
1488 fn advance(&mut self, cursor: &PadCopyCursor) {
1489 for axis in 0..cursor.shape.len() {
1493 let next = self.coords.as_mut_slice()[axis] + 1;
1494 if next < cursor.shape[axis] {
1495 self.coords.as_mut_slice()[axis] = next;
1496 self.source_offset += cursor.source_steps[axis];
1497 self.dest_offset += cursor.dest_steps[axis];
1498 return;
1499 }
1500 self.coords.as_mut_slice()[axis] = 0;
1501 self.source_offset += cursor.source_resets[axis];
1502 self.dest_offset += cursor.dest_resets[axis];
1503 }
1504 }
1505}
1506
1507struct CoordScratch {
1508 inline: [usize; crate::RAW_FUSED_RANK_LIMIT],
1509 heap: Option<Vec<usize>>,
1510 len: usize,
1511}
1512
1513impl CoordScratch {
1514 fn new(len: usize) -> Self {
1515 if len <= crate::RAW_FUSED_RANK_LIMIT {
1516 Self {
1517 inline: [0; crate::RAW_FUSED_RANK_LIMIT],
1518 heap: None,
1519 len,
1520 }
1521 } else {
1522 Self {
1523 inline: [0; crate::RAW_FUSED_RANK_LIMIT],
1524 heap: Some(vec![0; len]),
1525 len,
1526 }
1527 }
1528 }
1529
1530 fn as_mut_slice(&mut self) -> &mut [usize] {
1531 match &mut self.heap {
1532 Some(heap) => heap,
1533 None => &mut self.inline[..self.len],
1534 }
1535 }
1536}
1537
1538#[cfg(test)]
1539#[path = "static_indexing_plan/tests/tests.rs"]
1540mod tests;