use ndarray::prelude::*;
use serde::{Deserialize, Serialize};
#[derive(Clone, Debug, Eq, PartialEq, Serialize, Deserialize)]
pub struct MI {
shape: Array1<usize>,
strides: Array1<usize>,
}
impl MI {
pub fn new<I>(shape: I) -> Self
where
I: IntoIterator<Item = usize>,
{
let shape: Array1<_> = shape.into_iter().collect();
let mut strides: Vec<_> = shape
.iter()
.rev()
.scan(1, |acc, &dim| {
let stride = *acc;
*acc *= dim;
Some(stride)
})
.collect();
strides.reverse();
let strides = Array1::from(strides);
Self { shape, strides }
}
#[inline]
pub const fn shape(&self) -> &Array1<usize> {
&self.shape
}
pub fn ravel<I>(&self, multi_index: I) -> usize
where
I: IntoIterator<Item = usize>,
{
self.strides
.iter()
.zip(multi_index)
.map(|(i, j)| i * j)
.sum()
}
pub fn unravel(&self, index: usize) -> Vec<usize> {
let mut remaining_index = index;
self.strides
.iter()
.map(|&stride| {
let value = remaining_index / stride;
remaining_index %= stride;
value
})
.collect()
}
}