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