1use crate::{
2 checked_logical_element_count, checked_product, col_major_strides, validate_permutation,
3 DynRank, Error, Result, ShapeVec, SliceSpec, StrideVec, TensorRank,
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)).map_err(|_| Error::IntegerOverflow)?;
28 let delta = last.checked_mul(stride).ok_or(Error::IntegerOverflow)?;
29 if delta < 0 {
30 min = min.checked_add(delta).ok_or(Error::IntegerOverflow)?;
31 } else {
32 max = max.checked_add(delta).ok_or(Error::IntegerOverflow)?;
33 }
34 }
35 Ok(Some((min, max)))
36}
37
38pub(crate) fn validate_reachable_bounds(
39 shape: &[usize],
40 strides: &[isize],
41 offset: isize,
42 buffer_len: usize,
43) -> Result<()> {
44 if shape.len() != strides.len() {
45 return Err(Error::RankMismatch {
46 expected: shape.len(),
47 actual: strides.len(),
48 });
49 }
50
51 match reachable_offset_range(shape, strides, offset)? {
52 Some((min, max)) => {
53 if min < 0 {
54 return Err(Error::ViewOutOfBounds);
55 }
56 let max = usize::try_from(max).map_err(|_| Error::IntegerOverflow)?;
57 if max < buffer_len {
58 Ok(())
59 } else {
60 Err(Error::ViewOutOfBounds)
61 }
62 }
63 None => {
64 if offset < 0 {
65 return Err(Error::ViewOutOfBounds);
66 }
67 let offset = usize::try_from(offset).map_err(|_| Error::IntegerOverflow)?;
68 if offset <= buffer_len {
69 Ok(())
70 } else {
71 Err(Error::ViewOutOfBounds)
72 }
73 }
74 }
75}
76
77fn layout_from_vecs<R: TensorRank>(
78 shape: ShapeVec,
79 strides: StrideVec,
80 offset: isize,
81 buffer_len: usize,
82) -> Result<TensorLayout<R>> {
83 TensorLayout::from_parts(
84 R::shape_from_vec(shape)?,
85 R::strides_from_vec(strides)?,
86 offset,
87 buffer_len,
88 )
89}
90
91fn positive_ceil_div(numerator: isize, denominator: isize) -> Result<usize> {
92 if numerator < 0 || denominator <= 0 {
93 return Err(Error::IntegerOverflow);
94 }
95 let extent = if numerator == 0 {
96 0
97 } else {
98 1 + (numerator - 1) / denominator
99 };
100 usize::try_from(extent).map_err(|_| Error::IntegerOverflow)
101}
102
103fn normalize_slice(slice: SliceSpec, axis_len: usize) -> Result<(isize, usize)> {
104 if slice.step == 0 {
105 return Err(Error::InvalidSliceStep { step: slice.step });
106 }
107 if axis_len == 0 {
108 return Ok((0, 0));
109 }
110
111 let axis_len = isize::try_from(axis_len).map_err(|_| Error::IntegerOverflow)?;
112 if slice.step > 0 {
113 let start = if slice.start < 0 {
114 slice
115 .start
116 .checked_add(axis_len)
117 .ok_or(Error::IntegerOverflow)?
118 } else {
119 slice.start
120 };
121 let end = if slice.end < 0 {
122 slice
123 .end
124 .checked_add(axis_len)
125 .ok_or(Error::IntegerOverflow)?
126 } else {
127 slice.end
128 };
129 if start < 0 || start > axis_len || end < 0 || end > axis_len {
130 return Err(Error::InvalidSliceBounds {
131 start: slice.start,
132 end: slice.end,
133 axis_len: usize::try_from(axis_len).map_err(|_| Error::IntegerOverflow)?,
134 });
135 }
136 if start >= end {
137 return Ok((start, 0));
138 }
139 return Ok((start, positive_ceil_div(end - start, slice.step)?));
140 }
141
142 let start = if slice.start < 0 {
143 slice
144 .start
145 .checked_add(axis_len)
146 .ok_or(Error::IntegerOverflow)?
147 } else {
148 slice.start
149 };
150 let end = slice.end;
151 if start < 0 || start >= axis_len || end < -1 || end >= axis_len {
152 return Err(Error::InvalidSliceBounds {
153 start: slice.start,
154 end: slice.end,
155 axis_len: usize::try_from(axis_len).map_err(|_| Error::IntegerOverflow)?,
156 });
157 }
158 if start <= end {
159 return Ok((start, 0));
160 }
161 let step = slice.step.checked_neg().ok_or(Error::IntegerOverflow)?;
162 Ok((start, positive_ceil_div(start - end, step)?))
163}
164
165#[derive(Clone, Debug, PartialEq, Eq)]
178pub struct TensorLayout<R: TensorRank = DynRank> {
179 shape: R::Shape,
180 strides: R::Strides,
181 offset: isize,
182}
183
184impl<R: TensorRank> TensorLayout<R> {
185 pub fn compact(shape: R::Shape) -> Result<Self> {
197 let strides = R::strides_from_vec(col_major_strides(shape.as_ref())?)?;
198 Ok(Self {
199 shape,
200 strides,
201 offset: 0,
202 })
203 }
204
205 pub fn from_parts(
222 shape: R::Shape,
223 strides: R::Strides,
224 offset: isize,
225 buffer_len: usize,
226 ) -> Result<Self> {
227 checked_logical_element_count(shape.as_ref())?;
228 validate_reachable_bounds(shape.as_ref(), strides.as_ref(), offset, buffer_len)?;
229 Ok(Self {
230 shape,
231 strides,
232 offset,
233 })
234 }
235
236 pub fn shape(&self) -> &[usize] {
248 self.shape.as_ref()
249 }
250
251 pub fn strides(&self) -> &[isize] {
263 self.strides.as_ref()
264 }
265
266 pub fn offset(&self) -> isize {
278 self.offset
279 }
280
281 pub fn is_compact_col_major(&self) -> Result<bool> {
293 if self.shape().contains(&0) {
294 return Ok(true);
295 }
296
297 col_major_strides(self.shape()).map(|strides| strides.as_slice() == self.strides())
298 }
299
300 pub fn validate_mutable_no_overlap(&self) -> Result<()> {
317 if self.shape().contains(&0) {
318 return Ok(());
319 }
320
321 for (&extent, &stride) in self.shape().iter().zip(self.strides()) {
322 if extent > 1 && stride == 0 {
323 return Err(Error::OverlappingMutableLayout);
324 }
325 }
326
327 let element_count = checked_product(self.shape())?;
328
329 let mut axes = self
330 .shape()
331 .iter()
332 .zip(self.strides())
333 .filter(|&(&extent, _)| extent > 1)
334 .map(|(&extent, &stride)| (extent, stride.unsigned_abs()))
335 .collect::<SmallVec<[(usize, usize); 8]>>();
336 axes.sort_by_key(|&(_, stride)| stride);
337
338 let mut span = 0usize;
339 for (extent, stride) in axes {
340 if stride <= span {
341 return self.validate_mutable_no_overlap_exact_or_reject(element_count);
342 }
343 span = span
344 .checked_add(
345 (extent - 1)
346 .checked_mul(stride)
347 .ok_or(Error::IntegerOverflow)?,
348 )
349 .ok_or(Error::IntegerOverflow)?;
350 }
351
352 Ok(())
353 }
354
355 fn validate_mutable_no_overlap_exact_or_reject(&self, element_count: usize) -> Result<()> {
356 if element_count > MUTABLE_NO_OVERLAP_EXACT_ELEMENT_LIMIT {
357 return Err(Error::OverlappingMutableLayout);
358 }
359
360 let mut seen = HashSet::with_capacity(element_count);
361 let rank = self.shape().len();
362 let mut indices = vec![0usize; rank];
363
364 loop {
365 let mut physical_offset = self.offset;
366 for (&index, &stride) in indices.iter().zip(self.strides()) {
367 let index = isize::try_from(index).map_err(|_| Error::IntegerOverflow)?;
368 let delta = index.checked_mul(stride).ok_or(Error::IntegerOverflow)?;
369 physical_offset = physical_offset
370 .checked_add(delta)
371 .ok_or(Error::IntegerOverflow)?;
372 }
373
374 if !seen.insert(physical_offset) {
375 return Err(Error::OverlappingMutableLayout);
376 }
377
378 let mut axis = 0;
379 while axis < rank {
380 indices[axis] += 1;
381 if indices[axis] < self.shape()[axis] {
382 break;
383 }
384 indices[axis] = 0;
385 axis += 1;
386 }
387 if axis == rank {
388 return Ok(());
389 }
390 }
391 }
392
393 pub fn transpose_view(&self, axes: impl AsRef<[usize]>) -> Result<Self> {
407 let axes = axes.as_ref();
408 validate_permutation(self.shape().len(), axes)?;
409 let shape = axes
410 .iter()
411 .map(|&axis| self.shape()[axis])
412 .collect::<ShapeVec>();
413 let strides = axes
414 .iter()
415 .map(|&axis| self.strides()[axis])
416 .collect::<StrideVec>();
417 Ok(Self {
418 shape: R::shape_from_vec(shape)?,
419 strides: R::strides_from_vec(strides)?,
420 offset: self.offset,
421 })
422 }
423
424 pub fn slice_view(&self, spec: impl AsRef<[SliceSpec]>, buffer_len: usize) -> Result<Self> {
438 let spec = spec.as_ref();
439 if spec.len() != self.shape().len() {
440 return Err(Error::RankMismatch {
441 expected: self.shape().len(),
442 actual: spec.len(),
443 });
444 }
445
446 let mut shape = ShapeVec::new();
447 let mut strides = StrideVec::new();
448 let mut offset = self.offset;
449 for ((&axis_len, &stride), &slice) in self
450 .shape()
451 .iter()
452 .zip(self.strides().iter())
453 .zip(spec.iter())
454 {
455 let (start, extent) = normalize_slice(slice, axis_len)?;
456 let start_offset = start.checked_mul(stride).ok_or(Error::IntegerOverflow)?;
457 offset = offset
458 .checked_add(start_offset)
459 .ok_or(Error::IntegerOverflow)?;
460 shape.push(extent);
461 strides.push(
462 stride
463 .checked_mul(slice.step)
464 .ok_or(Error::IntegerOverflow)?,
465 );
466 }
467 layout_from_vecs(shape, strides, offset, buffer_len)
468 }
469
470 pub fn reshape_view_as<R2: TensorRank>(
484 &self,
485 shape: R2::Shape,
486 buffer_len: usize,
487 ) -> Result<TensorLayout<R2>> {
488 if !self.is_compact_col_major()? {
489 return Err(Error::NonContiguousViewAsSlice);
490 }
491 let from = checked_product(self.shape())?;
492 let to = checked_product(shape.as_ref())?;
493 if from != to {
494 return Err(Error::ReshapeElementCountMismatch { from, to });
495 }
496 let strides = R2::strides_from_vec(col_major_strides(shape.as_ref())?)?;
497 TensorLayout::from_parts(shape, strides, self.offset, buffer_len)
498 }
499
500 pub fn broadcast_in_dim_view<R2: TensorRank>(
514 &self,
515 shape: R2::Shape,
516 broadcast_dims: impl AsRef<[usize]>,
517 buffer_len: usize,
518 ) -> Result<TensorLayout<R2>> {
519 let broadcast_dims = broadcast_dims.as_ref();
520 if broadcast_dims.len() != self.shape().len() {
521 return Err(Error::RankMismatch {
522 expected: self.shape().len(),
523 actual: broadcast_dims.len(),
524 });
525 }
526
527 let output_rank = shape.as_ref().len();
528 let mut seen = vec![false; output_rank];
529 let mut strides = StrideVec::new();
530 strides.resize(output_rank, 0);
531 for (input_axis, &output_axis) in broadcast_dims.iter().enumerate() {
532 if output_axis >= output_rank {
533 return Err(Error::AxisOutOfBounds {
534 axis: output_axis,
535 rank: output_rank,
536 });
537 }
538 if seen[output_axis] {
539 return Err(Error::DuplicateAxis { axis: output_axis });
540 }
541 seen[output_axis] = true;
542
543 let input_extent = self.shape()[input_axis];
544 let output_extent = shape.as_ref()[output_axis];
545 if input_extent != output_extent && input_extent != 1 {
546 return Err(Error::ShapeDataLengthMismatch {
547 expected: input_extent,
548 actual: output_extent,
549 });
550 }
551 if input_extent == output_extent {
552 strides[output_axis] = self.strides()[input_axis];
553 }
554 }
555
556 TensorLayout::from_parts(
557 shape,
558 R2::strides_from_vec(strides)?,
559 self.offset,
560 buffer_len,
561 )
562 }
563}
564
565#[cfg(test)]
566mod tests {
567 use super::positive_ceil_div;
568 use crate::Error;
569 use std::panic::{catch_unwind, AssertUnwindSafe};
570
571 #[test]
572 fn positive_ceil_div_rejects_invalid_preconditions_without_panicking() {
573 for (numerator, denominator) in [(-1, 1), (1, 0), (1, -1)] {
574 let result = catch_unwind(AssertUnwindSafe(|| {
575 positive_ceil_div(numerator, denominator)
576 }));
577
578 assert!(
579 result.is_ok(),
580 "invalid positive_ceil_div inputs should return Err"
581 );
582 assert!(matches!(result.unwrap(), Err(Error::IntegerOverflow)));
583 }
584 }
585}