use crate::error::{Error, Result};
use crate::shape::Shape;
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct Layout {
shape: Shape,
strides: Vec<usize>,
offset: usize,
}
impl Layout {
pub fn contiguous(shape: Shape) -> Self {
let strides = shape.stride_contiguous();
Layout {
shape,
strides,
offset: 0,
}
}
pub fn new(shape: Shape, strides: Vec<usize>, offset: usize) -> Self {
Layout {
shape,
strides,
offset,
}
}
pub fn shape(&self) -> &Shape {
&self.shape
}
pub fn strides(&self) -> &[usize] {
&self.strides
}
pub fn offset(&self) -> usize {
self.offset
}
pub fn rank(&self) -> usize {
self.shape.rank()
}
pub fn dims(&self) -> &[usize] {
self.shape.dims()
}
pub fn elem_count(&self) -> usize {
self.shape.elem_count()
}
pub fn is_contiguous(&self) -> bool {
self.offset == 0 && self.strides == self.shape.stride_contiguous()
}
pub fn transpose(&self, dim0: usize, dim1: usize) -> Result<Layout> {
let rank = self.rank();
if dim0 >= rank || dim1 >= rank {
return Err(Error::DimOutOfRange {
dim: dim0.max(dim1),
rank,
});
}
let mut new_dims = self.shape.dims().to_vec();
let mut new_strides = self.strides.clone();
new_dims.swap(dim0, dim1);
new_strides.swap(dim0, dim1);
Ok(Layout::new(Shape::new(new_dims), new_strides, self.offset))
}
pub fn narrow(&self, dim: usize, start: usize, len: usize) -> Result<Layout> {
let rank = self.rank();
if dim >= rank {
return Err(Error::DimOutOfRange { dim, rank });
}
let dim_size = self.shape.dims()[dim];
if start + len > dim_size {
return Err(Error::NarrowOutOfBounds {
dim,
start,
len,
dim_size,
});
}
let mut new_dims = self.shape.dims().to_vec();
new_dims[dim] = len;
let new_offset = self.offset + start * self.strides[dim];
Ok(Layout::new(
Shape::new(new_dims),
self.strides.clone(),
new_offset,
))
}
pub fn flat_index(&self, index: &[usize]) -> usize {
let mut flat = self.offset;
for (i, &idx) in index.iter().enumerate() {
flat += idx * self.strides[i];
}
flat
}
pub fn strided_indices(&self) -> StridedIter {
StridedIter::new(self)
}
}
pub struct StridedIter {
current: Vec<usize>,
dims: Vec<usize>,
strides: Vec<usize>,
offset: usize,
remaining: usize,
started: bool,
}
impl StridedIter {
fn new(layout: &Layout) -> Self {
let rank = layout.rank();
StridedIter {
current: vec![0; rank],
dims: layout.dims().to_vec(),
strides: layout.strides().to_vec(),
offset: layout.offset(),
remaining: layout.elem_count(),
started: false,
}
}
fn flat_index(&self) -> usize {
let mut idx = self.offset;
for i in 0..self.current.len() {
idx += self.current[i] * self.strides[i];
}
idx
}
fn advance(&mut self) {
let rank = self.dims.len();
for i in (0..rank).rev() {
self.current[i] += 1;
if self.current[i] < self.dims[i] {
return;
}
self.current[i] = 0;
}
}
}
impl Iterator for StridedIter {
type Item = usize;
fn next(&mut self) -> Option<usize> {
if self.remaining == 0 {
return None;
}
if self.started {
self.advance();
}
self.started = true;
self.remaining -= 1;
Some(self.flat_index())
}
fn size_hint(&self) -> (usize, Option<usize>) {
(self.remaining, Some(self.remaining))
}
}
impl ExactSizeIterator for StridedIter {}
#[cfg(test)]
mod tests {
use super::*;
use crate::shape::Shape;
#[test]
fn test_contiguous_layout() {
let layout = Layout::contiguous(Shape::from((2, 3)));
assert!(layout.is_contiguous());
assert_eq!(layout.strides(), &[3, 1]);
assert_eq!(layout.offset(), 0);
}
#[test]
fn test_contiguous_indices() {
let layout = Layout::contiguous(Shape::from((2, 3)));
let indices: Vec<usize> = layout.strided_indices().collect();
assert_eq!(indices, vec![0, 1, 2, 3, 4, 5]);
}
#[test]
fn test_transpose_layout() {
let layout = Layout::contiguous(Shape::from((2, 3)));
let transposed = layout.transpose(0, 1).unwrap();
assert_eq!(transposed.dims(), &[3, 2]);
assert_eq!(transposed.strides(), &[1, 3]);
assert!(!transposed.is_contiguous());
}
#[test]
fn test_transpose_indices() {
let layout = Layout::contiguous(Shape::from((2, 3)));
let transposed = layout.transpose(0, 1).unwrap();
let indices: Vec<usize> = transposed.strided_indices().collect();
assert_eq!(indices, vec![0, 3, 1, 4, 2, 5]);
}
#[test]
fn test_narrow() {
let layout = Layout::contiguous(Shape::from((4, 6)));
let narrowed = layout.narrow(1, 2, 3).unwrap();
assert_eq!(narrowed.dims(), &[4, 3]);
assert_eq!(narrowed.offset(), 2);
assert_eq!(narrowed.strides(), &[6, 1]); }
#[test]
fn test_narrow_out_of_bounds() {
let layout = Layout::contiguous(Shape::from((4, 6)));
assert!(layout.narrow(1, 5, 3).is_err()); }
#[test]
fn test_flat_index() {
let layout = Layout::contiguous(Shape::from((2, 3, 4)));
assert_eq!(layout.flat_index(&[1, 2, 3]), 23);
assert_eq!(layout.flat_index(&[0, 0, 0]), 0);
}
}