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, Error, Result, ShapeVec, SliceSpec, StrideVec, TensorRank,
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)).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/// Storage-neutral tensor layout metadata.
166///
167/// # Examples
168///
169/// ```rust
170/// use tenferro_tensor_core::{Rank, TensorLayout};
171///
172/// let layout = TensorLayout::<Rank<2>>::compact([2, 3])?;
173/// assert_eq!(layout.shape(), &[2, 3]);
174/// assert_eq!(layout.strides(), &[1, 2]);
175/// # Ok::<(), tenferro_tensor_core::Error>(())
176/// ```
177#[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    /// Create a compact column-major layout with zero offset.
186    ///
187    /// # Examples
188    ///
189    /// ```rust
190    /// use tenferro_tensor_core::{Rank, TensorLayout};
191    ///
192    /// let layout = TensorLayout::<Rank<2>>::compact([2, 3])?;
193    /// assert_eq!(layout.strides(), &[1, 2]);
194    /// # Ok::<(), tenferro_tensor_core::Error>(())
195    /// ```
196    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    /// Create a layout from shape, strides, element offset, and backing buffer length.
206    ///
207    /// # Examples
208    ///
209    /// ```rust
210    /// use tenferro_tensor_core::{DynRank, TensorLayout};
211    ///
212    /// let layout = TensorLayout::<DynRank>::from_parts(
213    ///     vec![2, 3].into(),
214    ///     vec![1, 2].into(),
215    ///     0,
216    ///     6,
217    /// )?;
218    /// assert!(layout.is_compact_col_major()?);
219    /// # Ok::<(), tenferro_tensor_core::Error>(())
220    /// ```
221    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    /// Return the layout shape.
237    ///
238    /// # Examples
239    ///
240    /// ```rust
241    /// use tenferro_tensor_core::{Rank, TensorLayout};
242    ///
243    /// let layout = TensorLayout::<Rank<1>>::compact([4])?;
244    /// assert_eq!(layout.shape(), &[4]);
245    /// # Ok::<(), tenferro_tensor_core::Error>(())
246    /// ```
247    pub fn shape(&self) -> &[usize] {
248        self.shape.as_ref()
249    }
250
251    /// Return the layout strides in element units.
252    ///
253    /// # Examples
254    ///
255    /// ```rust
256    /// use tenferro_tensor_core::{Rank, TensorLayout};
257    ///
258    /// let layout = TensorLayout::<Rank<2>>::compact([2, 3])?;
259    /// assert_eq!(layout.strides(), &[1, 2]);
260    /// # Ok::<(), tenferro_tensor_core::Error>(())
261    /// ```
262    pub fn strides(&self) -> &[isize] {
263        self.strides.as_ref()
264    }
265
266    /// Return the layout element offset.
267    ///
268    /// # Examples
269    ///
270    /// ```rust
271    /// use tenferro_tensor_core::{DynRank, TensorLayout};
272    ///
273    /// let layout = TensorLayout::<DynRank>::from_parts(vec![3].into(), vec![1].into(), 2, 5)?;
274    /// assert_eq!(layout.offset(), 2);
275    /// # Ok::<(), tenferro_tensor_core::Error>(())
276    /// ```
277    pub fn offset(&self) -> isize {
278        self.offset
279    }
280
281    /// Return whether the layout has compact column-major strides.
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!(layout.is_compact_col_major()?);
290    /// # Ok::<(), tenferro_tensor_core::Error>(())
291    /// ```
292    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    /// Validate that the layout can be used for mutable access without aliasing.
301    ///
302    /// Empty logical views are accepted. Non-empty layouts are accepted when a
303    /// conservative stride-span proof succeeds, or when exact enumeration of a
304    /// small bounded view proves that all logical elements map to distinct
305    /// physical offsets.
306    ///
307    /// # Examples
308    ///
309    /// ```rust
310    /// use tenferro_tensor_core::{DynRank, TensorLayout};
311    ///
312    /// let layout = TensorLayout::<DynRank>::from_parts(vec![3].into(), vec![-1].into(), 2, 3)?;
313    /// layout.validate_mutable_no_overlap()?;
314    /// # Ok::<(), tenferro_tensor_core::Error>(())
315    /// ```
316    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    /// Return a metadata-only axis permutation of this layout.
394    ///
395    /// # Examples
396    ///
397    /// ```rust
398    /// use tenferro_tensor_core::{Rank, TensorLayout};
399    ///
400    /// let layout = TensorLayout::<Rank<2>>::compact([2, 3])?;
401    /// let transposed = layout.transpose_view([1, 0])?;
402    /// assert_eq!(transposed.shape(), &[3, 2]);
403    /// assert_eq!(transposed.strides(), &[2, 1]);
404    /// # Ok::<(), tenferro_tensor_core::Error>(())
405    /// ```
406    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    /// Return a metadata-only slice of this layout.
425    ///
426    /// # Examples
427    ///
428    /// ```rust
429    /// use tenferro_tensor_core::{Rank, SliceSpec, TensorLayout};
430    ///
431    /// let layout = TensorLayout::<Rank<1>>::compact([4])?;
432    /// let view = layout.slice_view([SliceSpec { start: 3, end: -1, step: -2 }], 4)?;
433    /// assert_eq!(view.shape(), &[2]);
434    /// assert_eq!(view.strides(), &[-2]);
435    /// # Ok::<(), tenferro_tensor_core::Error>(())
436    /// ```
437    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    /// Return a metadata-only reshape of this compact column-major layout.
471    ///
472    /// # Examples
473    ///
474    /// ```rust
475    /// use tenferro_tensor_core::{Rank, TensorLayout};
476    ///
477    /// let layout = TensorLayout::<Rank<2>>::compact([2, 3])?;
478    /// let reshaped = layout.reshape_view_as::<Rank<1>>([6], 6)?;
479    /// assert_eq!(reshaped.shape(), &[6]);
480    /// assert_eq!(reshaped.strides(), &[1]);
481    /// # Ok::<(), tenferro_tensor_core::Error>(())
482    /// ```
483    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    /// Return a metadata-only explicit broadcast of this layout into a target rank.
501    ///
502    /// # Examples
503    ///
504    /// ```rust
505    /// use tenferro_tensor_core::{Rank, TensorLayout};
506    ///
507    /// let layout = TensorLayout::<Rank<1>>::compact([3])?;
508    /// let broadcast = layout.broadcast_in_dim_view::<Rank<2>>([2, 3], [1], 3)?;
509    /// assert_eq!(broadcast.shape(), &[2, 3]);
510    /// assert_eq!(broadcast.strides(), &[0, 1]);
511    /// # Ok::<(), tenferro_tensor_core::Error>(())
512    /// ```
513    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}