1use crate::{
2 checked_logical_element_count, checked_product, col_major_strides, validate_permutation,
3 DynRank, Result, ShapeMismatch, ShapeVec, SliceSpec, StrideVec, TensorRank, ValidationError,
4};
5use smallvec::SmallVec;
6use std::collections::HashSet;
7
8const MUTABLE_NO_OVERLAP_EXACT_ELEMENT_LIMIT: usize = 4096;
14
15pub(crate) fn reachable_offset_range(
16 shape: &[usize],
17 strides: &[isize],
18 offset: isize,
19) -> Result<Option<(isize, isize)>> {
20 if shape.contains(&0) {
21 return Ok(None);
22 }
23
24 let mut min = offset;
25 let mut max = offset;
26 for (&extent, &stride) in shape.iter().zip(strides) {
27 let last = isize::try_from(extent.saturating_sub(1))
28 .map_err(|_| ValidationError::IntegerOverflow)?;
29 let delta = last
30 .checked_mul(stride)
31 .ok_or(ValidationError::IntegerOverflow)?;
32 if delta < 0 {
33 min = min
34 .checked_add(delta)
35 .ok_or(ValidationError::IntegerOverflow)?;
36 } else {
37 max = max
38 .checked_add(delta)
39 .ok_or(ValidationError::IntegerOverflow)?;
40 }
41 }
42 Ok(Some((min, max)))
43}
44
45pub(crate) fn validate_reachable_bounds(
46 shape: &[usize],
47 strides: &[isize],
48 offset: isize,
49 buffer_len: usize,
50) -> Result<()> {
51 if shape.len() != strides.len() {
52 return Err(ValidationError::RankMismatch {
53 expected: shape.len(),
54 actual: strides.len(),
55 });
56 }
57
58 match reachable_offset_range(shape, strides, offset)? {
59 Some((min, max)) => {
60 if min < 0 {
61 return Err(ValidationError::ViewOutOfBounds);
62 }
63 let max = usize::try_from(max).map_err(|_| ValidationError::IntegerOverflow)?;
64 if max < buffer_len {
65 Ok(())
66 } else {
67 Err(ValidationError::ViewOutOfBounds)
68 }
69 }
70 None => {
71 if offset < 0 {
72 return Err(ValidationError::ViewOutOfBounds);
73 }
74 let offset = usize::try_from(offset).map_err(|_| ValidationError::IntegerOverflow)?;
75 if offset <= buffer_len {
76 Ok(())
77 } else {
78 Err(ValidationError::ViewOutOfBounds)
79 }
80 }
81 }
82}
83
84fn layout_from_vecs<R: TensorRank>(
85 shape: ShapeVec,
86 strides: StrideVec,
87 offset: isize,
88 buffer_len: usize,
89) -> Result<TensorLayout<R>> {
90 TensorLayout::from_parts(
91 R::shape_from_vec(shape)?,
92 R::strides_from_vec(strides)?,
93 offset,
94 buffer_len,
95 )
96}
97
98fn positive_ceil_div(numerator: isize, denominator: isize) -> Result<usize> {
99 if numerator < 0 || denominator <= 0 {
100 return Err(ValidationError::IntegerOverflow);
101 }
102 let extent = if numerator == 0 {
103 0
104 } else {
105 1 + (numerator - 1) / denominator
106 };
107 usize::try_from(extent).map_err(|_| ValidationError::IntegerOverflow)
108}
109
110fn normalize_slice(slice: SliceSpec, axis_len: usize) -> Result<(isize, usize)> {
111 if slice.step == 0 {
112 return Err(ValidationError::InvalidSliceStep { step: slice.step });
113 }
114 if axis_len == 0 {
115 return Ok((0, 0));
116 }
117
118 let axis_len = isize::try_from(axis_len).map_err(|_| ValidationError::IntegerOverflow)?;
119 if slice.step > 0 {
120 let start = if slice.start < 0 {
121 slice
122 .start
123 .checked_add(axis_len)
124 .ok_or(ValidationError::IntegerOverflow)?
125 } else {
126 slice.start
127 };
128 let end = if slice.end < 0 {
129 slice
130 .end
131 .checked_add(axis_len)
132 .ok_or(ValidationError::IntegerOverflow)?
133 } else {
134 slice.end
135 };
136 if start < 0 || start > axis_len || end < 0 || end > axis_len {
137 return Err(ValidationError::InvalidSliceBounds {
138 start: slice.start,
139 end: slice.end,
140 axis_len: usize::try_from(axis_len)
141 .map_err(|_| ValidationError::IntegerOverflow)?,
142 });
143 }
144 if start >= end {
145 return Ok((start, 0));
146 }
147 return Ok((start, positive_ceil_div(end - start, slice.step)?));
148 }
149
150 let start = if slice.start < 0 {
151 slice
152 .start
153 .checked_add(axis_len)
154 .ok_or(ValidationError::IntegerOverflow)?
155 } else {
156 slice.start
157 };
158 let end = if slice.end < -1 {
159 slice
160 .end
161 .checked_add(axis_len)
162 .ok_or(ValidationError::IntegerOverflow)?
163 } else {
164 slice.end
165 };
166 if start < 0 || start >= axis_len || end < -1 || end >= axis_len {
167 return Err(ValidationError::InvalidSliceBounds {
168 start: slice.start,
169 end: slice.end,
170 axis_len: usize::try_from(axis_len).map_err(|_| ValidationError::IntegerOverflow)?,
171 });
172 }
173 if start <= end {
174 return Ok((start, 0));
175 }
176 let step = slice
177 .step
178 .checked_neg()
179 .ok_or(ValidationError::IntegerOverflow)?;
180 Ok((start, positive_ceil_div(start - end, step)?))
181}
182
183#[derive(Clone, Debug, PartialEq, Eq)]
196pub struct TensorLayout<R: TensorRank = DynRank> {
197 shape: R::Shape,
198 strides: R::Strides,
199 offset: isize,
200}
201
202impl<R: TensorRank> TensorLayout<R> {
203 pub fn compact(shape: R::Shape) -> Result<Self> {
220 let strides = R::strides_from_vec(col_major_strides(shape.as_ref())?)?;
221 Ok(Self {
222 shape,
223 strides,
224 offset: 0,
225 })
226 }
227
228 pub fn from_parts(
252 shape: R::Shape,
253 strides: R::Strides,
254 offset: isize,
255 buffer_len: usize,
256 ) -> Result<Self> {
257 checked_logical_element_count(shape.as_ref())?;
258 validate_reachable_bounds(shape.as_ref(), strides.as_ref(), offset, buffer_len)?;
259 Ok(Self {
260 shape,
261 strides,
262 offset,
263 })
264 }
265
266 pub fn shape(&self) -> &[usize] {
278 self.shape.as_ref()
279 }
280
281 pub fn strides(&self) -> &[isize] {
293 self.strides.as_ref()
294 }
295
296 pub fn offset(&self) -> isize {
308 self.offset
309 }
310
311 pub fn is_compact_col_major(&self) -> Result<bool> {
328 if self.shape().contains(&0) {
329 return Ok(true);
330 }
331
332 col_major_strides(self.shape()).map(|strides| strides.as_slice() == self.strides())
333 }
334
335 pub fn validate_mutable_no_overlap(&self) -> Result<()> {
358 if self.shape().contains(&0) {
359 return Ok(());
360 }
361
362 for (&extent, &stride) in self.shape().iter().zip(self.strides()) {
363 if extent > 1 && stride == 0 {
364 return Err(ValidationError::OverlappingMutableLayout);
365 }
366 }
367
368 let element_count = checked_product(self.shape())?;
369
370 let mut axes = self
371 .shape()
372 .iter()
373 .zip(self.strides())
374 .filter(|&(&extent, _)| extent > 1)
375 .map(|(&extent, &stride)| (extent, stride.unsigned_abs()))
376 .collect::<SmallVec<[(usize, usize); 8]>>();
377 axes.sort_by_key(|&(_, stride)| stride);
378
379 let mut span = 0usize;
380 for (extent, stride) in axes {
381 if stride <= span {
382 return self.validate_mutable_no_overlap_exact_or_reject(element_count);
383 }
384 span = span
385 .checked_add(
386 (extent - 1)
387 .checked_mul(stride)
388 .ok_or(ValidationError::IntegerOverflow)?,
389 )
390 .ok_or(ValidationError::IntegerOverflow)?;
391 }
392
393 Ok(())
394 }
395
396 fn validate_mutable_no_overlap_exact_or_reject(&self, element_count: usize) -> Result<()> {
397 if element_count > MUTABLE_NO_OVERLAP_EXACT_ELEMENT_LIMIT {
398 return Err(ValidationError::OverlappingMutableLayout);
399 }
400
401 let mut seen = HashSet::with_capacity(element_count);
402 let rank = self.shape().len();
403 let mut indices = vec![0usize; rank];
404
405 loop {
406 let mut physical_offset = self.offset;
407 for (&index, &stride) in indices.iter().zip(self.strides()) {
408 let index = isize::try_from(index).map_err(|_| ValidationError::IntegerOverflow)?;
409 let delta = index
410 .checked_mul(stride)
411 .ok_or(ValidationError::IntegerOverflow)?;
412 physical_offset = physical_offset
413 .checked_add(delta)
414 .ok_or(ValidationError::IntegerOverflow)?;
415 }
416
417 if !seen.insert(physical_offset) {
418 return Err(ValidationError::OverlappingMutableLayout);
419 }
420
421 let mut axis = 0;
422 while axis < rank {
423 indices[axis] += 1;
424 if indices[axis] < self.shape()[axis] {
425 break;
426 }
427 indices[axis] = 0;
428 axis += 1;
429 }
430 if axis == rank {
431 return Ok(());
432 }
433 }
434 }
435
436 pub fn transpose_view(&self, axes: impl AsRef<[usize]>) -> Result<Self> {
457 let axes = axes.as_ref();
458 validate_permutation(self.shape().len(), axes)?;
459 let shape = axes
460 .iter()
461 .map(|&axis| self.shape()[axis])
462 .collect::<ShapeVec>();
463 let strides = axes
464 .iter()
465 .map(|&axis| self.strides()[axis])
466 .collect::<StrideVec>();
467 Ok(Self {
468 shape: R::shape_from_vec(shape)?,
469 strides: R::strides_from_vec(strides)?,
470 offset: self.offset,
471 })
472 }
473
474 pub fn slice_view(&self, spec: impl AsRef<[SliceSpec]>, buffer_len: usize) -> Result<Self> {
496 let spec = spec.as_ref();
497 if spec.len() != self.shape().len() {
498 return Err(ValidationError::RankMismatch {
499 expected: self.shape().len(),
500 actual: spec.len(),
501 });
502 }
503
504 let mut shape = ShapeVec::new();
505 let mut strides = StrideVec::new();
506 let mut offset = self.offset;
507 for ((&axis_len, &stride), &slice) in self
508 .shape()
509 .iter()
510 .zip(self.strides().iter())
511 .zip(spec.iter())
512 {
513 let (start, extent) = normalize_slice(slice, axis_len)?;
514 let start_offset = start
515 .checked_mul(stride)
516 .ok_or(ValidationError::IntegerOverflow)?;
517 offset = offset
518 .checked_add(start_offset)
519 .ok_or(ValidationError::IntegerOverflow)?;
520 shape.push(extent);
521 strides.push(
522 stride
523 .checked_mul(slice.step)
524 .ok_or(ValidationError::IntegerOverflow)?,
525 );
526 }
527 layout_from_vecs(shape, strides, offset, buffer_len)
528 }
529
530 pub fn reshape_view_as<R2: TensorRank>(
551 &self,
552 shape: impl Into<R2::Shape>,
553 buffer_len: usize,
554 ) -> Result<TensorLayout<R2>> {
555 let shape = shape.into();
556 if !self.is_compact_col_major()? {
557 return Err(ValidationError::NonContiguousViewAsSlice);
558 }
559 let from = checked_product(self.shape())?;
560 let to = checked_product(shape.as_ref())?;
561 if from != to {
562 return Err(ShapeMismatch::ReshapeElementCount { from, to }.into());
563 }
564 let strides = R2::strides_from_vec(col_major_strides(shape.as_ref())?)?;
565 TensorLayout::from_parts(shape, strides, self.offset, buffer_len)
566 }
567
568 pub fn broadcast_in_dim_view<R2: TensorRank>(
590 &self,
591 shape: impl Into<R2::Shape>,
592 broadcast_dims: impl AsRef<[usize]>,
593 buffer_len: usize,
594 ) -> Result<TensorLayout<R2>> {
595 let shape = shape.into();
596 let broadcast_dims = broadcast_dims.as_ref();
597 if broadcast_dims.len() != self.shape().len() {
598 return Err(ValidationError::RankMismatch {
599 expected: self.shape().len(),
600 actual: broadcast_dims.len(),
601 });
602 }
603
604 let output_rank = shape.as_ref().len();
605 let mut seen = vec![false; output_rank];
606 let mut strides = StrideVec::new();
607 strides.resize(output_rank, 0);
608 for (input_axis, &output_axis) in broadcast_dims.iter().enumerate() {
609 if output_axis >= output_rank {
610 return Err(ValidationError::AxisOutOfBounds {
611 axis: output_axis,
612 rank: output_rank,
613 });
614 }
615 if seen[output_axis] {
616 return Err(ValidationError::DuplicateAxis {
617 axis: output_axis,
618 role: "permutation",
619 });
620 }
621 seen[output_axis] = true;
622
623 let input_extent = self.shape()[input_axis];
624 let output_extent = shape.as_ref()[output_axis];
625 if input_extent != output_extent && input_extent != 1 {
626 return Err(ValidationError::ShapeDataLengthMismatch {
627 expected: input_extent,
628 actual: output_extent,
629 });
630 }
631 if input_extent == output_extent {
632 strides[output_axis] = self.strides()[input_axis];
633 }
634 }
635
636 TensorLayout::from_parts(
637 shape,
638 R2::strides_from_vec(strides)?,
639 self.offset,
640 buffer_len,
641 )
642 }
643}
644
645#[cfg(test)]
646mod tests {
647 use super::positive_ceil_div;
648 use crate::ValidationError;
649 use std::panic::{catch_unwind, AssertUnwindSafe};
650
651 #[test]
652 fn positive_ceil_div_rejects_invalid_preconditions_without_panicking() {
653 for (numerator, denominator) in [(-1, 1), (1, 0), (1, -1)] {
654 let result = catch_unwind(AssertUnwindSafe(|| {
655 positive_ceil_div(numerator, denominator)
656 }));
657
658 assert!(
659 result.is_ok(),
660 "invalid positive_ceil_div inputs should return Err"
661 );
662 assert!(matches!(
663 result.unwrap(),
664 Err(ValidationError::IntegerOverflow)
665 ));
666 }
667 }
668}