pub fn is_injective_layout(dims: &[usize], strides: &[isize]) -> bool {
let Some(total) = validate_injective_layout_inputs(dims, strides) else {
return false;
};
if total <= 1 {
return true;
}
match interleaved_block(dims, strides) {
None => false,
Some(None) => true,
Some(Some(block)) => {
block.may_be_injective()
&& block.total <= EXACT_BLOCK_BUDGET as u128
&& block_offsets_unique(dims, strides, &block)
}
}
}
pub(crate) fn is_injective_layout_without_alloc(dims: &[usize], strides: &[isize]) -> bool {
let Some(total) = validate_injective_layout_inputs(dims, strides) else {
return false;
};
if total <= 1 {
return true;
}
match interleaved_block(dims, strides) {
None => false,
Some(None) => true,
Some(Some(block)) => {
block.may_be_injective()
&& block.total <= PAIRWISE_BLOCK_BUDGET as u128
&& block_offsets_unique_pairwise(dims, strides, &block)
}
}
}
pub(crate) const EXACT_BLOCK_BUDGET: usize = 1 << 24;
pub(crate) const PAIRWISE_BLOCK_BUDGET: usize = 4096;
struct InterleavedBlock {
last_key: (u128, usize),
total: u128,
span: u128,
}
impl InterleavedBlock {
fn contains(&self, axis: usize, dim: usize, stride: isize) -> bool {
dim > 1 && (stride.unsigned_abs() as u128, axis) <= self.last_key
}
fn may_be_injective(&self) -> bool {
self.total <= self.span + 1
}
}
fn interleaved_block(dims: &[usize], strides: &[isize]) -> Option<Option<InterleavedBlock>> {
let mut covered_span = 0u128;
let mut covered_total = 1u128;
let mut previous_key = None;
let mut block = None;
let active_axes = dims.iter().filter(|&&dim| dim > 1).count();
for _ in 0..active_axes {
let mut next = None;
for (axis, (&dim, &stride)) in dims.iter().zip(strides.iter()).enumerate() {
if dim <= 1 {
continue;
}
let key = (stride.checked_abs()? as u128, axis);
if previous_key.is_some_and(|previous| key <= previous) {
continue;
}
if next.is_none_or(|(best, _)| key < best) {
next = Some((key, dim as u128));
}
}
let ((stride, axis), dim) = next?;
let separated = stride > covered_span;
covered_span = stride
.checked_mul(dim - 1)
.and_then(|span| covered_span.checked_add(span))?;
covered_total = covered_total.checked_mul(dim)?;
if !separated {
block = Some(InterleavedBlock {
last_key: (stride, axis),
total: covered_total,
span: covered_span,
});
}
previous_key = Some((stride, axis));
}
Some(block)
}
fn block_offsets_unique(dims: &[usize], strides: &[isize], block: &InterleavedBlock) -> bool {
let axes: Vec<(usize, usize)> = dims
.iter()
.zip(strides.iter())
.enumerate()
.filter(|&(axis, (&dim, &stride))| block.contains(axis, dim, stride))
.map(|(_, (&dim, &stride))| (dim, stride.unsigned_abs()))
.collect();
let (Ok(total), Ok(span)) = (usize::try_from(block.total), usize::try_from(block.span)) else {
return false;
};
let mut indices = vec![0usize; axes.len()];
let mut offset = 0usize;
let mut advance = |offset: &mut usize| {
for (index, &(dim, stride)) in indices.iter_mut().zip(axes.iter()) {
if *index + 1 < dim {
*index += 1;
*offset += stride;
return;
}
*offset -= stride * (dim - 1);
*index = 0;
}
};
if block.span / 64 <= block.total {
let mut seen = vec![0u64; span / 64 + 1];
for _ in 0..total {
let (word, bit) = (offset / 64, 1u64 << (offset % 64));
if seen[word] & bit != 0 {
return false;
}
seen[word] |= bit;
advance(&mut offset);
}
true
} else {
let mut offsets = Vec::with_capacity(total);
for _ in 0..total {
offsets.push(offset);
advance(&mut offset);
}
offsets.sort_unstable();
offsets.windows(2).all(|pair| pair[0] != pair[1])
}
}
fn block_offset_for_linear_index(
dims: &[usize],
strides: &[isize],
block: &InterleavedBlock,
mut linear: usize,
) -> u128 {
let mut offset = 0u128;
for (axis, (&dim, &stride)) in dims.iter().zip(strides.iter()).enumerate() {
if !block.contains(axis, dim, stride) {
continue;
}
offset += stride.unsigned_abs() as u128 * (linear % dim) as u128;
linear /= dim;
}
offset
}
fn block_offsets_unique_pairwise(
dims: &[usize],
strides: &[isize],
block: &InterleavedBlock,
) -> bool {
let Ok(total) = usize::try_from(block.total) else {
return false;
};
for lhs in 0..total {
let lhs_offset = block_offset_for_linear_index(dims, strides, block, lhs);
for rhs in (lhs + 1)..total {
if block_offset_for_linear_index(dims, strides, block, rhs) == lhs_offset {
return false;
}
}
}
true
}
fn validate_injective_layout_inputs(dims: &[usize], strides: &[isize]) -> Option<usize> {
if dims.len() != strides.len() {
return None;
}
let total = crate::kernel::total_len(dims).ok()?;
if total <= 1 {
return Some(total);
}
if dims
.iter()
.zip(strides.iter())
.any(|(&dim, &stride)| dim > 1 && stride == 0)
{
return None;
}
let mut min_offset = 0isize;
let mut max_offset = 0isize;
for (&dim, &stride) in dims.iter().zip(strides.iter()) {
if dim <= 1 {
continue;
}
let extent = isize::try_from(dim - 1).ok()?;
let span = stride.checked_mul(extent)?;
if span >= 0 {
max_offset = max_offset.checked_add(span)?;
} else {
min_offset = min_offset.checked_add(span)?;
}
}
Some(total)
}