use crate::dtype::DataType;
use crate::error::IrError;
pub fn compute_contiguous_strides(shape: &[usize]) -> Vec<i64> {
let n = shape.len();
let mut strides = vec![1i64; n];
for i in (0..n.saturating_sub(1)).rev() {
strides[i] = strides[i + 1] * shape[i + 1] as i64;
}
strides
}
pub fn is_contiguous(shape: &[usize], strides: &[i64]) -> bool {
if shape.len() != strides.len() {
return false;
}
let mut expected: i64 = 1;
for i in (0..shape.len()).rev() {
if strides[i] != expected {
return false;
}
expected *= shape[i] as i64;
}
true
}
pub fn is_dense(shape: &[usize], strides: &[i64]) -> bool {
if shape.len() != strides.len() {
return false;
}
let mut expected: i64 = 1;
let mut fast = true;
for i in (0..shape.len()).rev() {
if shape[i] == 0 || strides[i] != expected {
fast = false;
break;
}
expected *= shape[i] as i64;
}
if fast {
return true;
}
const INLINE_RANK: usize = 8;
let nontrivial = |(&d, &s): (&usize, &i64)| (s.unsigned_abs() as i64, d);
if shape.len() <= INLINE_RANK {
let mut pairs = [(0i64, 0usize); INLINE_RANK];
let mut len = 0;
for pair in shape
.iter()
.zip(strides)
.filter(|&(&d, _)| d > 1)
.map(nontrivial)
{
pairs[len] = pair;
len += 1;
}
dense_extents(&mut pairs[..len])
} else {
let mut pairs: Vec<(i64, usize)> = shape
.iter()
.zip(strides)
.filter(|&(&d, _)| d > 1)
.map(nontrivial)
.collect();
dense_extents(&mut pairs)
}
}
fn dense_extents(pairs: &mut [(i64, usize)]) -> bool {
if pairs.is_empty() {
return true; }
pairs.sort_unstable_by_key(|&(s, _)| s);
if pairs[0].0 != 1 {
return false;
}
let mut expected_stride: i64 = 1;
for &(stride, size) in &*pairs {
if stride != expected_stride {
return false;
}
expected_stride *= size as i64;
}
true
}
pub fn broadcast_shapes(a: &[usize], b: &[usize]) -> Result<Vec<usize>, IrError> {
let max_ndim = a.len().max(b.len());
let mut result = Vec::with_capacity(max_ndim);
for i in 0..max_ndim {
let da = if i < a.len() { a[a.len() - 1 - i] } else { 1 };
let db = if i < b.len() { b[b.len() - 1 - i] } else { 1 };
if da == db || db == 1 {
result.push(da);
} else if da == 1 {
result.push(db);
} else {
return Err(IrError::BroadcastIncompatible {
a: a.to_vec(),
b: b.to_vec(),
});
}
}
result.reverse();
Ok(result)
}
#[derive(Clone, Debug, PartialEq, Eq, Default)]
pub enum MemoryFormat {
#[default]
Contiguous,
ChannelsLast,
Blocked(usize),
Custom,
}
#[derive(Clone, Debug, PartialEq)]
pub struct TensorLayout {
pub strides: Option<Vec<i64>>,
pub format: MemoryFormat,
pub alignment: usize,
}
pub const DEFAULT_ALIGNMENT: usize = 64;
impl Default for TensorLayout {
fn default() -> Self {
Self {
strides: None,
format: MemoryFormat::Contiguous,
alignment: DEFAULT_ALIGNMENT,
}
}
}
impl TensorLayout {
pub fn contiguous() -> Self {
Self::default()
}
pub fn strided(strides: Vec<i64>) -> Self {
Self {
strides: Some(strides),
format: MemoryFormat::Custom,
alignment: DEFAULT_ALIGNMENT,
}
}
pub fn is_contiguous(&self, shape: &[usize]) -> bool {
match &self.strides {
None => true,
Some(s) => is_contiguous(shape, s),
}
}
pub fn resolved_strides(&self, shape: &[usize]) -> Vec<i64> {
self.strides
.clone()
.unwrap_or_else(|| compute_contiguous_strides(shape))
}
pub fn transpose(&self, shape: &[usize], perm: &[usize]) -> Self {
let base = self.resolved_strides(shape);
let strides = perm.iter().map(|&p| base[p]).collect();
Self {
strides: Some(strides),
format: MemoryFormat::Custom,
alignment: self.alignment,
}
}
pub fn storage_size(&self, shape: &[usize], dtype: DataType) -> usize {
let elem = dtype.byte_size().max(1);
match &self.strides {
None => shape.iter().product::<usize>() * elem,
Some(strides) => {
let max_offset: i64 = shape
.iter()
.zip(strides.iter())
.map(|(&dim, &stride)| dim.saturating_sub(1) as i64 * stride.abs())
.sum();
(max_offset as usize + 1) * elem
}
}
}
}
#[cfg(test)]
mod tests {
use super::*;
fn is_contiguous_reference(shape: &[usize], strides: &[i64]) -> bool {
strides == compute_contiguous_strides(shape).as_slice()
}
#[test]
fn contiguous_walk_agrees_with_the_materialising_implementation() {
let mut state: u64 = 0x9E37_79B9_7F4A_7C15;
let mut next = |bound: u64| -> u64 {
state = state
.wrapping_mul(6_364_136_223_846_793_005)
.wrapping_add(1_442_695_040_888_963_407);
(state >> 33) % bound
};
let mut agreed_true = 0usize;
for rank in 0..=6usize {
for _ in 0..500 {
let shape: Vec<usize> = (0..rank).map(|_| next(4) as usize).collect();
let strides: Vec<i64> = match next(3) {
0 => compute_contiguous_strides(&shape),
1 => {
let mut s = compute_contiguous_strides(&shape);
if !s.is_empty() {
let i = next(s.len() as u64) as usize;
s[i] += next(3) as i64 - 1;
}
s
}
_ => (0..rank).map(|_| next(9) as i64 - 4).collect(),
};
let got = is_contiguous(&shape, &strides);
if got && !shape.is_empty() {
agreed_true += 1;
}
assert_eq!(
got,
is_contiguous_reference(&shape, &strides),
"disagreement for shape {shape:?} strides {strides:?}"
);
let mut long = strides.clone();
long.push(1);
assert_eq!(
is_contiguous(&shape, &long),
is_contiguous_reference(&shape, &long),
"disagreement for shape {shape:?} strides {long:?}"
);
}
}
assert!(
agreed_true > 100,
"the corpus never reached the accepting arm on a non-empty shape \
(only {agreed_true} cases), so it proved nothing"
);
}
fn is_dense_reference(shape: &[usize], strides: &[i64]) -> bool {
if shape.len() != strides.len() {
return false;
}
let mut pairs: Vec<(i64, usize)> = shape
.iter()
.zip(strides)
.filter(|&(&d, _)| d > 1)
.map(|(&d, &s)| (s.unsigned_abs() as i64, d))
.collect();
if pairs.is_empty() {
return true;
}
pairs.sort_unstable_by_key(|&(s, _)| s);
if pairs[0].0 != 1 {
return false;
}
let mut expected_stride: i64 = 1;
for &(stride, size) in &pairs {
if stride != expected_stride {
return false;
}
expected_stride *= size as i64;
}
true
}
#[test]
fn the_contiguity_shortcut_rejects_zero_extents() {
assert!(is_contiguous(&[2, 0], &[0, 1]), "premise of this test");
assert!(!is_dense(&[2, 0], &[0, 1]));
assert_eq!(
is_dense(&[2, 0], &[0, 1]),
is_dense_reference(&[2, 0], &[0, 1])
);
assert!(is_dense(&[0], &[1]));
assert!(is_dense(&[0, 3], &[3, 1]));
}
#[test]
fn the_contiguity_shortcut_does_not_shadow_the_general_path() {
assert!(!is_contiguous(&[4, 3], &[1, 4]));
assert!(is_dense(&[4, 3], &[1, 4]));
assert!(is_dense(&[2, 3], &[3, -1]));
assert_eq!(
is_dense(&[2, 3], &[3, -1]),
is_dense_reference(&[2, 3], &[3, -1])
);
let shape = [2usize, 1, 2, 1, 2, 1, 2, 1, 2, 2];
let strides = compute_contiguous_strides(&shape);
assert!(is_dense(&shape, &strides));
assert_eq!(
is_dense(&shape, &strides),
is_dense_reference(&shape, &strides)
);
}
#[test]
fn inline_storage_agrees_with_the_original_implementation() {
let mut state: u64 = 0x2545_F491_4F6C_DD1D;
let mut next = |bound: u64| -> u64 {
state = state
.wrapping_mul(6_364_136_223_846_793_005)
.wrapping_add(1_442_695_040_888_963_407);
(state >> 33) % bound
};
let mut checked_dense = 0usize;
for rank in 0..=10usize {
for _ in 0..400 {
let shape: Vec<usize> = (0..rank).map(|_| next(4) as usize).collect();
let strides: Vec<i64> = if next(2) == 0 {
compute_contiguous_strides(&shape)
} else {
(0..rank)
.map(|_| next(9) as i64 - 4) .collect()
};
let got = is_dense(&shape, &strides);
if got && shape.iter().any(|&d| d > 1) {
checked_dense += 1;
}
assert_eq!(
got,
is_dense_reference(&shape, &strides),
"disagreement for shape {shape:?} strides {strides:?}"
);
let mut short = strides.clone();
short.pop();
assert_eq!(
is_dense(&shape, &short),
is_dense_reference(&shape, &short),
"disagreement for shape {shape:?} strides {short:?}"
);
}
}
assert!(
checked_dense > 100,
"the corpus degenerated into rejections and trivial shapes; it proved \
nothing about the sort-and-product arm (only {checked_dense} dense \
cases with a dimension above 1)"
);
}
#[test]
fn a_completely_full_inline_array_is_handled() {
let shape = [2usize; 8];
assert_eq!(shape.len(), 8, "this test must fill the inline array");
let strides = compute_contiguous_strides(&shape);
assert!(is_dense(&shape, &strides));
assert!(is_dense_reference(&shape, &strides));
let mut permuted: Vec<i64> = strides.clone();
permuted.reverse();
assert!(!is_contiguous(&shape, &permuted));
assert_eq!(
is_dense(&shape, &permuted),
is_dense_reference(&shape, &permuted)
);
assert!(is_dense(&shape, &permuted));
let mut holed = strides.clone();
holed[0] += 1;
assert!(!is_dense(&shape, &holed));
assert!(!is_dense_reference(&shape, &holed));
}
#[test]
fn ranks_above_the_inline_bound_use_the_heap_path_correctly() {
let shape = [2usize, 2, 2, 2, 2, 2, 2, 2, 2];
let strides = compute_contiguous_strides(&shape);
assert!(shape.len() > 8, "this test must exercise the fallback");
assert!(is_dense(&shape, &strides));
assert!(is_dense_reference(&shape, &strides));
let mut broken = strides.clone();
broken[0] += 1;
assert!(!is_dense(&shape, &broken));
assert!(!is_dense_reference(&shape, &broken));
}
#[test]
fn contiguous_strides_row_major() {
assert_eq!(compute_contiguous_strides(&[2, 3, 4]), vec![12, 4, 1]);
assert_eq!(compute_contiguous_strides(&[5]), vec![1]);
assert_eq!(compute_contiguous_strides(&[]), Vec::<i64>::new());
}
#[test]
fn is_contiguous_check() {
assert!(is_contiguous(&[2, 3], &[3, 1]));
assert!(!is_contiguous(&[2, 3], &[1, 2]));
}
#[test]
fn broadcast_basic() {
assert_eq!(broadcast_shapes(&[3, 1], &[1, 4]).unwrap(), vec![3, 4]);
assert_eq!(broadcast_shapes(&[5], &[3, 5]).unwrap(), vec![3, 5]);
assert_eq!(broadcast_shapes(&[], &[2, 2]).unwrap(), vec![2, 2]);
}
#[test]
fn broadcast_incompatible() {
assert!(matches!(
broadcast_shapes(&[3], &[4]),
Err(IrError::BroadcastIncompatible { .. })
));
}
#[test]
fn transpose_swaps_strides() {
let l = TensorLayout::contiguous();
let t = l.transpose(&[2, 3], &[1, 0]);
assert_eq!(t.strides, Some(vec![1, 3]));
assert!(!t.is_contiguous(&[3, 2]));
}
#[test]
fn storage_size_contiguous_and_strided() {
let l = TensorLayout::contiguous();
assert_eq!(l.storage_size(&[2, 3], DataType::Float32), 24);
let t = l.transpose(&[2, 3], &[1, 0]);
assert_eq!(t.storage_size(&[3, 2], DataType::Float32), 24);
}
}