Skip to main content

tenferro_tensor_core/
layout.rs

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
9/// Maximum logical elements for exact mutable-overlap validation.
10///
11/// Larger layouts must pass the sufficient stride-span proof. This keeps the
12/// fallback bounded because it enumerates logical elements and stores visited
13/// physical offsets.
14const 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/// Storage-neutral tensor layout metadata.
185///
186/// # Examples
187///
188/// ```rust
189/// use tenferro_tensor_core::{Rank, TensorLayout};
190///
191/// let layout = TensorLayout::<Rank<2>>::compact([2, 3])?;
192/// assert_eq!(layout.shape(), &[2, 3]);
193/// assert_eq!(layout.strides(), &[1, 2]);
194/// # Ok::<(), tenferro_tensor_core::ValidationError>(())
195/// ```
196#[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    /// Create a compact column-major layout with zero offset.
205    ///
206    /// # Examples
207    ///
208    /// ```rust
209    /// use tenferro_tensor_core::{Rank, TensorLayout};
210    ///
211    /// let layout = TensorLayout::<Rank<2>>::compact([2, 3])?;
212    /// assert_eq!(layout.strides(), &[1, 2]);
213    /// # Ok::<(), tenferro_tensor_core::ValidationError>(())
214    /// ```
215    ///
216    /// # Errors
217    ///
218    /// Returns [`ValidationError::IntegerOverflow`] when compact strides
219    /// cannot be computed for `shape`.
220    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    /// Create a layout from shape, strides, element offset, and backing buffer length.
230    ///
231    /// # Examples
232    ///
233    /// ```rust
234    /// use tenferro_tensor_core::{DynRank, TensorLayout};
235    ///
236    /// let layout = TensorLayout::<DynRank>::from_parts(
237    ///     vec![2, 3].into(),
238    ///     vec![1, 2].into(),
239    ///     0,
240    ///     6,
241    /// )?;
242    /// assert!(layout.is_compact_col_major()?);
243    /// # Ok::<(), tenferro_tensor_core::ValidationError>(())
244    /// ```
245    ///
246    /// # Errors
247    ///
248    /// Returns [`ValidationError::RankMismatch`] for incompatible shape and
249    /// stride ranks, [`ValidationError::ViewOutOfBounds`] when the reachable
250    /// range exceeds `buffer_len`, or [`ValidationError::IntegerOverflow`]
251    /// when metadata arithmetic overflows.
252    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    /// Return the layout shape.
268    ///
269    /// # Examples
270    ///
271    /// ```rust
272    /// use tenferro_tensor_core::{Rank, TensorLayout};
273    ///
274    /// let layout = TensorLayout::<Rank<1>>::compact([4])?;
275    /// assert_eq!(layout.shape(), &[4]);
276    /// # Ok::<(), tenferro_tensor_core::ValidationError>(())
277    /// ```
278    pub fn shape(&self) -> &[usize] {
279        self.shape.as_ref()
280    }
281
282    /// Return the layout strides in element units.
283    ///
284    /// # Examples
285    ///
286    /// ```rust
287    /// use tenferro_tensor_core::{Rank, TensorLayout};
288    ///
289    /// let layout = TensorLayout::<Rank<2>>::compact([2, 3])?;
290    /// assert_eq!(layout.strides(), &[1, 2]);
291    /// # Ok::<(), tenferro_tensor_core::ValidationError>(())
292    /// ```
293    pub fn strides(&self) -> &[isize] {
294        self.strides.as_ref()
295    }
296
297    /// Return the layout element offset.
298    ///
299    /// # Examples
300    ///
301    /// ```rust
302    /// use tenferro_tensor_core::{DynRank, TensorLayout};
303    ///
304    /// let layout = TensorLayout::<DynRank>::from_parts(vec![3].into(), vec![1].into(), 2, 5)?;
305    /// assert_eq!(layout.offset(), 2);
306    /// # Ok::<(), tenferro_tensor_core::ValidationError>(())
307    /// ```
308    pub fn offset(&self) -> isize {
309        self.offset
310    }
311
312    /// Return whether the layout has compact column-major strides.
313    ///
314    /// # Examples
315    ///
316    /// ```rust
317    /// use tenferro_tensor_core::{Rank, TensorLayout};
318    ///
319    /// let layout = TensorLayout::<Rank<2>>::compact([2, 3])?;
320    /// assert!(layout.is_compact_col_major()?);
321    /// # Ok::<(), tenferro_tensor_core::ValidationError>(())
322    /// ```
323    ///
324    /// # Errors
325    ///
326    /// Returns [`ValidationError::IntegerOverflow`] when compactness
327    /// validation overflows metadata arithmetic.
328    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    /// Validate that the layout can be used for mutable access without aliasing.
337    ///
338    /// Empty logical views are accepted. Non-empty layouts are accepted when a
339    /// conservative stride-span proof succeeds, or when exact enumeration of a
340    /// small bounded view proves that all logical elements map to distinct
341    /// physical offsets.
342    ///
343    /// # Examples
344    ///
345    /// ```rust
346    /// use tenferro_tensor_core::{DynRank, TensorLayout};
347    ///
348    /// let layout = TensorLayout::<DynRank>::from_parts(vec![3].into(), vec![-1].into(), 2, 3)?;
349    /// layout.validate_mutable_no_overlap()?;
350    /// # Ok::<(), tenferro_tensor_core::ValidationError>(())
351    /// ```
352    ///
353    /// # Errors
354    ///
355    /// Returns [`ValidationError::OverlappingMutableLayout`] when multiple
356    /// logical elements can alias, or [`ValidationError::IntegerOverflow`]
357    /// when overlap validation arithmetic overflows.
358    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    /// Return a metadata-only axis permutation of this layout.
438    ///
439    /// # Examples
440    ///
441    /// ```rust
442    /// use tenferro_tensor_core::{Rank, TensorLayout};
443    ///
444    /// let layout = TensorLayout::<Rank<2>>::compact([2, 3])?;
445    /// let transposed = layout.transpose_view([1, 0])?;
446    /// assert_eq!(transposed.shape(), &[3, 2]);
447    /// assert_eq!(transposed.strides(), &[2, 1]);
448    /// # Ok::<(), tenferro_tensor_core::ValidationError>(())
449    /// ```
450    ///
451    /// # Errors
452    ///
453    /// Returns [`ValidationError::InvalidPermutationLength`],
454    /// [`ValidationError::AxisOutOfBounds`], or
455    /// [`ValidationError::DuplicateAxis`] when `axes` is not a permutation of
456    /// the layout rank.
457    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    /// Return a metadata-only slice of this layout.
476    ///
477    /// # Examples
478    ///
479    /// ```rust
480    /// use tenferro_tensor_core::{Rank, SliceSpec, TensorLayout};
481    ///
482    /// let layout = TensorLayout::<Rank<1>>::compact([4])?;
483    /// let view = layout.slice_view([SliceSpec { start: 3, end: -1, step: -2 }], 4)?;
484    /// assert_eq!(view.shape(), &[2]);
485    /// assert_eq!(view.strides(), &[-2]);
486    /// # Ok::<(), tenferro_tensor_core::ValidationError>(())
487    /// ```
488    ///
489    /// # Errors
490    ///
491    /// Returns [`ValidationError::RankMismatch`] when `spec` does not cover
492    /// every axis, [`ValidationError::InvalidSliceStep`] or
493    /// [`ValidationError::InvalidSliceBounds`] for invalid slice parameters,
494    /// or [`ValidationError::ViewOutOfBounds`] when the result is outside the
495    /// backing buffer.
496    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    /// Return a metadata-only reshape of this compact column-major layout.
532    ///
533    /// # Examples
534    ///
535    /// ```rust
536    /// use tenferro_tensor_core::{Rank, TensorLayout};
537    ///
538    /// let layout = TensorLayout::<Rank<2>>::compact([2, 3])?;
539    /// let reshaped = layout.reshape_view_as::<Rank<1>>([6], 6)?;
540    /// assert_eq!(reshaped.shape(), &[6]);
541    /// assert_eq!(reshaped.strides(), &[1]);
542    /// # Ok::<(), tenferro_tensor_core::ValidationError>(())
543    /// ```
544    ///
545    /// # Errors
546    ///
547    /// Returns [`ValidationError::NonContiguousViewAsSlice`] for a noncompact
548    /// source, [`ValidationError::ShapeMismatch`] for a different element
549    /// count, or [`ValidationError::IntegerOverflow`] when shape arithmetic
550    /// overflows.
551    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    /// Return a metadata-only explicit broadcast of this layout into a target rank.
570    ///
571    /// # Examples
572    ///
573    /// ```rust
574    /// use tenferro_tensor_core::{Rank, TensorLayout};
575    ///
576    /// let layout = TensorLayout::<Rank<1>>::compact([3])?;
577    /// let broadcast = layout.broadcast_in_dim_view::<Rank<2>>([2, 3], [1], 3)?;
578    /// assert_eq!(broadcast.shape(), &[2, 3]);
579    /// assert_eq!(broadcast.strides(), &[0, 1]);
580    /// # Ok::<(), tenferro_tensor_core::ValidationError>(())
581    /// ```
582    ///
583    /// # Errors
584    ///
585    /// Returns [`ValidationError::RankMismatch`],
586    /// [`ValidationError::AxisOutOfBounds`], or
587    /// [`ValidationError::DuplicateAxis`] for invalid broadcast axes;
588    /// [`ValidationError::ShapeDataLengthMismatch`] for incompatible extents;
589    /// or [`ValidationError::ViewOutOfBounds`] for an invalid result layout.
590    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    /// Return the compile-time-rank shape without erasing its array type.
648    ///
649    /// # Examples
650    ///
651    /// ```rust
652    /// use tenferro_tensor_core::{Rank, TensorLayout};
653    ///
654    /// let layout = TensorLayout::<Rank<2>>::compact([2, 3])?;
655    /// assert_eq!(layout.shape_array(), &[2, 3]);
656    /// # Ok::<(), tenferro_tensor_core::ValidationError>(())
657    /// ```
658    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}