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