Skip to main content

luma_tensor/tensor/
shape.rs

1use super::{Dim, DimCoordinates, DimNCoordinates};
2use crate::{Error, Result};
3use std::vec;
4
5#[derive(Debug, Clone, PartialEq, Eq)]
6pub struct Shape(pub(crate) Vec<usize>);
7
8impl Shape {
9    pub fn scalar() -> Self {
10        Self(vec![])
11    }
12
13    pub fn is_scalar(&self) -> bool {
14        self.0.is_empty() || (self.0.len() == 1 && self.0[0] == 1)
15    }
16
17    pub fn rank(&self) -> usize {
18        self.0.len()
19    }
20
21    pub fn dims(&self) -> &[usize] {
22        &self.0
23    }
24
25    pub fn into_dims(self) -> Vec<usize> {
26        self.0
27    }
28
29    pub fn dim(&self, dim: impl Dim) -> Result<usize> {
30        let index = dim.to_index(self, "get dim")?;
31        Ok(self.dims()[index])
32    }
33
34    pub fn element_count(&self) -> usize {
35        self.dims().iter().product()
36    }
37
38    pub fn is_contiguous(&self, stride: &[usize]) -> bool {
39        if self.rank() != stride.len() {
40            return false;
41        }
42        // [3, 4, 5] & [20, 5, 1]
43        let mut acc = 1;
44        for (&stride, &dim) in stride.iter().zip(self.dims().iter()).rev() {
45            if dim > 1 && stride != acc {
46                return false;
47            }
48            acc *= dim;
49        }
50        true
51    }
52
53    pub fn extend(mut self, additional_dims: &[usize]) -> Self {
54        self.0.extend(additional_dims);
55        self
56    }
57
58    /// Check whether the two shapes are compatible for broadcast, and if it is the case return the
59    /// broadcasted shape. This is to be used for binary pointwise ops.
60    /// Copy from https://github.com/huggingface/candle
61    pub fn broadcast_shape_binary_op(&self, rhs: &Self, op: &'static str) -> Result<Shape> {
62        let lhs = self;
63        let lhs_dims = lhs.dims();
64        let rhs_dims = rhs.dims();
65        let lhs_ndims = lhs_dims.len();
66        let rhs_ndims = rhs_dims.len();
67        let bcast_ndims = usize::max(lhs_ndims, rhs_ndims);
68        let mut bcast_dims = vec![0; bcast_ndims];
69        for (idx, bcast_value) in bcast_dims.iter_mut().enumerate() {
70            let rev_idx = bcast_ndims - idx;
71            let l_value = if lhs_ndims < rev_idx { 1 } else { lhs_dims[lhs_ndims - rev_idx] };
72            let r_value = if rhs_ndims < rev_idx { 1 } else { rhs_dims[rhs_ndims - rev_idx] };
73            *bcast_value = if l_value == r_value {
74                // keep
75                l_value
76            } else if l_value == 1 {
77                // bcast l
78                r_value
79            } else if r_value == 1 {
80                // bcast r
81                l_value
82            } else {
83                Err(Error::ShapeMismatchBinaryOp { lhs: lhs.clone(), rhs: rhs.clone(), op })?
84            }
85        }
86        Ok(Shape::from(bcast_dims))
87    }
88
89    /// Returns an iterator over **dimension coordinates**.
90    ///
91    /// This iterator yields the multi-dimensional coordinates
92    /// (e.g., `[i, j, k, ...]`) of each element in the array, independent
93    /// of the physical storage layout.
94    ///
95    /// Example for shape = (2, 2):
96    /// yields: `[0, 0], [0, 1], [1, 0], [1, 1]`
97    pub fn dim_coordinates(&self) -> DimCoordinates {
98        DimCoordinates::from_shape(self)
99    }
100
101    pub fn dims_coordinates<const N: usize>(&self) -> Result<DimNCoordinates<N>> {
102        DimNCoordinates::<N>::from_shape(self)
103    }
104
105    pub fn dim2_coordinates(&self) -> Result<DimNCoordinates<2>> {
106        DimNCoordinates::<2>::from_shape(self)
107    }
108
109    pub fn dim3_coordinates(&self) -> Result<DimNCoordinates<3>> {
110        DimNCoordinates::<3>::from_shape(self)
111    }
112
113    pub fn dim4_coordinates(&self) -> Result<DimNCoordinates<4>> {
114        DimNCoordinates::<4>::from_shape(self)
115    }
116
117    pub fn dim5_coordinates(&self) -> Result<DimNCoordinates<5>> {
118        DimNCoordinates::<5>::from_shape(self)
119    }
120
121    pub(crate) fn stride_contiguous(&self) -> Vec<usize> {
122        let mut stride = self
123            .dims()
124            .iter()
125            .rev()
126            .scan(1, |prod, u| {
127                let prod_pre_mult = *prod;
128                *prod *= u;
129                Some(prod_pre_mult)
130            })
131            .collect::<Vec<_>>();
132        stride.reverse();
133        stride
134    }
135}