luma_tensor/tensor/
shape.rs1use 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 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 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 l_value
76 } else if l_value == 1 {
77 r_value
79 } else if r_value == 1 {
80 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 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}