pub fn shape_size(shape: &Vec<usize>) -> usize {
let mut size: usize = 1;
for dim_size in shape.iter() {
size = size * dim_size;
}
size
}
pub fn plain_index(shape: &Vec<usize>, index_shape: &Vec<usize>, index: &Vec<usize>) -> usize {
let shape_len = shape.len();
assert_eq!(shape_len, index_shape.len());
assert_eq!(shape_len, index.len());
let mut plain_index: usize = 0;
for i in 0..shape_len {
if shape[i] == 1 {
continue;
}
if shape[i] != index_shape[i] {
panic!(
"Can't calculate equivalent plain index for shape {:#?} to {:#?}",
shape, index_shape
);
}
let mut mult: usize = 1;
for j in (i + 1)..shape_len {
mult *= shape[j];
}
plain_index += index[i] * mult
}
plain_index
}
pub fn next_index(shape: &Vec<usize>, index: &mut Vec<usize>) -> bool {
let shape_len = shape.len();
assert_eq!(shape_len, index.len());
for i in (0..shape_len).rev() {
index[i] += 1;
if index[i] == shape[i] {
index[i] = 0
} else {
return false;
}
}
true
}
pub fn dim_expansion_for_shape(shape: &Vec<usize>, n_dim: usize) -> Vec<usize> {
assert!(shape.len() <= n_dim);
let mut new_shape = vec![1; n_dim];
for i in 0..shape.len() {
new_shape[n_dim - shape.len() + i] = shape[i]
}
new_shape
}
pub fn adaptative_result_shape(lhs_shape: &Vec<usize>, rhs_shape: &Vec<usize>) -> Vec<usize> {
let lhs_shape_len = lhs_shape.len();
let rhs_shape_len = rhs_shape.len();
let out_shape_len = if lhs_shape_len > rhs_shape_len {
lhs_shape_len
} else {
rhs_shape_len
};
let mut out_shape = vec![1_usize; out_shape_len];
for i in 0..out_shape_len {
let lhs_dim_size = if i >= lhs_shape_len {
1_usize
} else {
lhs_shape[lhs_shape_len - 1 - i]
};
let rhs_dim_size = if i >= rhs_shape_len {
1_usize
} else {
rhs_shape[rhs_shape_len - 1 - i]
};
if lhs_dim_size == rhs_dim_size || rhs_dim_size == 1 {
out_shape[out_shape_len - 1 - i] = lhs_dim_size;
} else if lhs_dim_size == 1 {
out_shape[out_shape_len - 1 - i] = rhs_dim_size;
} else {
panic!("Incompatible shapes for perform commutative operation");
}
}
out_shape
}
pub trait Shape {
fn shape(&self) -> Vec<usize>;
fn shape_size(&self) -> usize {
shape_size(&self.shape())
}
fn plain_index(&self, index: &Vec<usize>) -> usize {
let shape = self.shape();
plain_index(&shape, &shape, index)
}
}
pub trait Reshape: Shape {
fn reshape(&mut self, new_shape: &Vec<usize>);
}