use alloc::vec::Vec;
use crate::Shape;
pub type DimOrder = Shape;
pub fn dim_order(shape: &[usize], strides: &[usize]) -> Option<DimOrder> {
dim_order_inner(shape, strides, Padding::Rejected)
}
pub fn nested_dim_order(shape: &[usize], strides: &[usize]) -> Option<DimOrder> {
dim_order_inner(shape, strides, Padding::Allowed)
}
#[derive(Clone, Copy)]
enum Padding {
Allowed,
Rejected,
}
fn dim_order_inner(shape: &[usize], strides: &[usize], padding: Padding) -> Option<DimOrder> {
let rank = shape.len();
if rank != strides.len() {
return None;
}
let mut order: Vec<usize> = (0..rank).collect();
order.sort_by(|a, b| strides[*b].cmp(&strides[*a]).then(a.cmp(b)));
let mut expected = 1;
for &axis in order.iter().rev() {
if shape[axis] == 1 {
continue;
}
match padding {
Padding::Allowed if strides[axis] < expected => return None,
Padding::Rejected if strides[axis] != expected => return None,
_ => {}
}
expected = strides[axis] * shape[axis];
}
Some(Shape::from(order))
}
pub fn is_contiguous_order(order: &[usize]) -> bool {
order.iter().enumerate().all(|(pos, axis)| pos == *axis)
}
#[cfg(test)]
mod tests {
use super::*;
use alloc::vec;
#[test]
fn contiguous_is_the_identity_order() {
let shape = [2, 48, 16, 16];
let strides = [48 * 16 * 16, 16 * 16, 16, 1];
assert_eq!(
dim_order(&shape, &strides),
Some(Shape::from(vec![0, 1, 2, 3]))
);
}
#[test]
fn nhwc_memory_gives_the_nhwc_order() {
let shape = [2, 48, 16, 16];
let strides = [16 * 16 * 48, 1, 16 * 48, 48];
assert_eq!(
dim_order(&shape, &strides),
Some(Shape::from(vec![0, 2, 3, 1]))
);
}
#[test]
fn broadcast_is_not_dense() {
let shape = [2, 48, 16, 16];
let strides = [0, 1, 0, 0];
assert_eq!(dim_order(&shape, &strides), None);
}
#[test]
fn a_slice_is_not_dense() {
let shape = [4, 8];
let strides = [16, 1];
assert_eq!(dim_order(&shape, &strides), None);
}
#[test]
fn size_one_dimensions_do_not_decide_density() {
let shape = [1, 48, 1, 1];
let strides = [48, 1, 48, 48];
assert!(dim_order(&shape, &strides).is_some());
}
#[test]
fn the_order_ends_at_the_innermost_dimension() {
let shape = [2, 48, 16, 16];
let contiguous = dim_order(&shape, &[48 * 16 * 16, 16 * 16, 16, 1]).unwrap();
let nhwc = dim_order(&shape, &[16 * 16 * 48, 1, 16 * 48, 48]).unwrap();
assert_eq!(contiguous.last(), Some(&3));
assert_eq!(nhwc.last(), Some(&1));
}
#[test]
fn order_is_a_permutation() {
assert!(is_contiguous_order(&[0, 1, 2, 3]));
assert!(!is_contiguous_order(&[0, 2, 3, 1]));
}
#[test]
fn a_dense_tensor_nests_in_the_order_it_is_dense_in() {
let shape = [2, 48, 16, 16];
for strides in [
[48 * 16 * 16, 16 * 16, 16, 1],
[16 * 16 * 48, 1, 16 * 48, 48],
] {
let dense = dim_order(&shape, &strides);
assert!(dense.is_some());
assert_eq!(nested_dim_order(&shape, &strides), dense);
}
}
#[test]
fn padding_under_the_innermost_dimension_keeps_the_nhwc_order() {
let shape = [2, 48, 16, 16];
let strides = [16 * 16 * 64, 1, 16 * 64, 64];
assert_eq!(dim_order(&shape, &strides), None);
assert_eq!(
nested_dim_order(&shape, &strides),
Some(Shape::from(vec![0, 2, 3, 1]))
);
}
#[test]
fn padding_above_the_innermost_dimension_is_nesting_too() {
let shape = [4, 8];
let strides = [16, 1];
assert_eq!(dim_order(&shape, &strides), None);
assert_eq!(
nested_dim_order(&shape, &strides),
Some(Shape::from(vec![0, 1]))
);
}
#[test]
fn overlapping_dimensions_are_not_an_order_under_either() {
let shape = [4, 8];
let strides = [4, 1];
assert_eq!(dim_order(&shape, &strides), None);
assert_eq!(nested_dim_order(&shape, &strides), None);
}
#[test]
fn a_broadcast_dimension_still_cannot_vote() {
let shape = [2, 48, 16, 16];
let strides = [0, 1, 0, 0];
assert_eq!(nested_dim_order(&shape, &strides), None);
}
#[test]
fn size_one_dimensions_neither_pad_nor_constrain_nesting() {
let shape = [1, 48, 1, 1];
let strides = [48, 1, 48, 48];
assert_eq!(
nested_dim_order(&shape, &strides),
dim_order(&shape, &strides)
);
assert!(nested_dim_order(&shape, &strides).is_some());
}
}