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 = Array1::from_elem(shape.len(), 1);
for i in (0..shape.len().saturating_sub(1)).rev() {
strides[i] = strides[i + 1] * shape[i + 1];
}
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 multi_index = Vec::with_capacity(self.shape.len());
let mut remaining_index = index;
for &stride in &self.strides {
let value = remaining_index / stride;
multi_index.push(value);
remaining_index -= value * stride;
}
multi_index
}
}