Skip to main content

strided_basic/
layout_check.rs

1//! Output layout validation shared by kernel families.
2
3/// Decide whether distinct logical output positions map to distinct offsets.
4///
5/// Axes are analysed in ascending `|stride|` order. An axis whose stride
6/// exceeds the offset span already covered by all smaller-stride axes is
7/// separated: it can never alias them, so it does not need enumeration.
8/// Only the interleaved block of smaller-stride axes that precede the last
9/// non-separated axis is checked exactly, by enumerating its offsets once
10/// with incremental traversal. The answer is exact whenever that block holds
11/// at most [`EXACT_BLOCK_BUDGET`] logical elements, independent of the total
12/// destination size (issues #255 and #256). Larger interleaved blocks, and
13/// metadata whose spans cannot be represented, are conservatively rejected.
14///
15/// # Examples
16///
17/// ```
18/// use strided_basic::execution::is_injective_layout;
19/// assert!(is_injective_layout(&[2, 3], &[1, 2]));
20/// assert!(!is_injective_layout(&[2, 3], &[0, 1]));
21/// // Interleaved but injective: 2000 = 2 (mod 3) keeps the rows apart.
22/// assert!(is_injective_layout(&[3, 1500], &[2000, 3]));
23/// ```
24pub fn is_injective_layout(dims: &[usize], strides: &[isize]) -> bool {
25    let Some(total) = validate_injective_layout_inputs(dims, strides) else {
26        return false;
27    };
28    if total <= 1 {
29        return true;
30    }
31    match interleaved_block(dims, strides) {
32        None => false,
33        Some(None) => true,
34        Some(Some(block)) => {
35            block.may_be_injective()
36                && block.total <= EXACT_BLOCK_BUDGET as u128
37                && block_offsets_unique(dims, strides, &block)
38        }
39    }
40}
41
42/// Allocation-free variant of [`is_injective_layout`].
43///
44/// Uses the same separated-axis reduction, then compares the interleaved
45/// block pairwise without allocating. It is exact for interleaved blocks of
46/// at most [`PAIRWISE_BLOCK_BUDGET`] logical elements and conservatively
47/// rejects larger ones.
48pub(crate) fn is_injective_layout_without_alloc(dims: &[usize], strides: &[isize]) -> bool {
49    let Some(total) = validate_injective_layout_inputs(dims, strides) else {
50        return false;
51    };
52    if total <= 1 {
53        return true;
54    }
55    match interleaved_block(dims, strides) {
56        None => false,
57        Some(None) => true,
58        Some(Some(block)) => {
59            block.may_be_injective()
60                && block.total <= PAIRWISE_BLOCK_BUDGET as u128
61                && block_offsets_unique_pairwise(dims, strides, &block)
62        }
63    }
64}
65
66/// Largest interleaved block enumerated exactly by [`is_injective_layout`].
67///
68/// INVARIANT: bounds the auxiliary memory of the exact check. The block is
69/// visited once in O(block) time; its seen-set is a bitmap over the block's
70/// offset span when that span is at most 64 offsets per element, otherwise a
71/// sorted offset list, so it never exceeds 8 bytes per block element
72/// (128 MiB at this bound). Separated axes never count toward the budget.
73pub(crate) const EXACT_BLOCK_BUDGET: usize = 1 << 24;
74
75/// Largest interleaved block compared pairwise by the allocation-free check.
76///
77/// INVARIANT: keeps the O(block^2) pairwise comparison of
78/// `is_injective_layout_without_alloc` bounded; separated axes never count
79/// toward the budget.
80pub(crate) const PAIRWISE_BLOCK_BUDGET: usize = 4096;
81
82/// The smallest-stride axes that must be checked by enumeration.
83///
84/// Contains every non-singleton axis whose `(|stride|, axis)` key is at most
85/// `last_key`. All remaining axes are separated from this block.
86struct InterleavedBlock {
87    last_key: (u128, usize),
88    /// Product of the block's extents.
89    total: u128,
90    /// Largest offset reachable inside the block after normalising strides to
91    /// their absolute values (the smallest is zero).
92    span: u128,
93}
94
95impl InterleavedBlock {
96    fn contains(&self, axis: usize, dim: usize, stride: isize) -> bool {
97        dim > 1 && (stride.unsigned_abs() as u128, axis) <= self.last_key
98    }
99
100    /// Pigeonhole: more elements than reachable offsets always alias.
101    fn may_be_injective(&self) -> bool {
102        self.total <= self.span + 1
103    }
104}
105
106/// Find the interleaved block of a layout.
107///
108/// Returns `None` when a stride or span cannot be represented, `Some(None)`
109/// when every axis is separated (so the layout is injective), and the block
110/// otherwise. Negating a stride only reflects the coordinate along that axis
111/// and shifts all offsets by a constant, so the analysis uses `|stride|`.
112fn interleaved_block(dims: &[usize], strides: &[isize]) -> Option<Option<InterleavedBlock>> {
113    let mut covered_span = 0u128;
114    let mut covered_total = 1u128;
115    let mut previous_key = None;
116    let mut block = None;
117    let active_axes = dims.iter().filter(|&&dim| dim > 1).count();
118    for _ in 0..active_axes {
119        let mut next = None;
120        for (axis, (&dim, &stride)) in dims.iter().zip(strides.iter()).enumerate() {
121            if dim <= 1 {
122                continue;
123            }
124            let key = (stride.checked_abs()? as u128, axis);
125            if previous_key.is_some_and(|previous| key <= previous) {
126                continue;
127            }
128            if next.is_none_or(|(best, _)| key < best) {
129                next = Some((key, dim as u128));
130            }
131        }
132        let ((stride, axis), dim) = next?;
133        let separated = stride > covered_span;
134        covered_span = stride
135            .checked_mul(dim - 1)
136            .and_then(|span| covered_span.checked_add(span))?;
137        covered_total = covered_total.checked_mul(dim)?;
138        if !separated {
139            block = Some(InterleavedBlock {
140                last_key: (stride, axis),
141                total: covered_total,
142                span: covered_span,
143            });
144        }
145        previous_key = Some((stride, axis));
146    }
147    Some(block)
148}
149
150/// Enumerate the block's offsets once and report whether they are distinct.
151fn block_offsets_unique(dims: &[usize], strides: &[isize], block: &InterleavedBlock) -> bool {
152    let axes: Vec<(usize, usize)> = dims
153        .iter()
154        .zip(strides.iter())
155        .enumerate()
156        .filter(|&(axis, (&dim, &stride))| block.contains(axis, dim, stride))
157        .map(|(_, (&dim, &stride))| (dim, stride.unsigned_abs()))
158        .collect();
159    let (Ok(total), Ok(span)) = (usize::try_from(block.total), usize::try_from(block.span)) else {
160        return false;
161    };
162
163    // Every visited offset and every partial axis span lies within
164    // [0, span], which fits `usize` (checked above), so the incremental
165    // updates below cannot overflow.
166    let mut indices = vec![0usize; axes.len()];
167    let mut offset = 0usize;
168    let mut advance = |offset: &mut usize| {
169        for (index, &(dim, stride)) in indices.iter_mut().zip(axes.iter()) {
170            if *index + 1 < dim {
171                *index += 1;
172                *offset += stride;
173                return;
174            }
175            *offset -= stride * (dim - 1);
176            *index = 0;
177        }
178    };
179
180    if block.span / 64 <= block.total {
181        let mut seen = vec![0u64; span / 64 + 1];
182        for _ in 0..total {
183            let (word, bit) = (offset / 64, 1u64 << (offset % 64));
184            if seen[word] & bit != 0 {
185                return false;
186            }
187            seen[word] |= bit;
188            advance(&mut offset);
189        }
190        true
191    } else {
192        let mut offsets = Vec::with_capacity(total);
193        for _ in 0..total {
194            offsets.push(offset);
195            advance(&mut offset);
196        }
197        offsets.sort_unstable();
198        offsets.windows(2).all(|pair| pair[0] != pair[1])
199    }
200}
201
202fn block_offset_for_linear_index(
203    dims: &[usize],
204    strides: &[isize],
205    block: &InterleavedBlock,
206    mut linear: usize,
207) -> u128 {
208    let mut offset = 0u128;
209    for (axis, (&dim, &stride)) in dims.iter().zip(strides.iter()).enumerate() {
210        if !block.contains(axis, dim, stride) {
211            continue;
212        }
213        offset += stride.unsigned_abs() as u128 * (linear % dim) as u128;
214        linear /= dim;
215    }
216    offset
217}
218
219fn block_offsets_unique_pairwise(
220    dims: &[usize],
221    strides: &[isize],
222    block: &InterleavedBlock,
223) -> bool {
224    let Ok(total) = usize::try_from(block.total) else {
225        return false;
226    };
227    for lhs in 0..total {
228        let lhs_offset = block_offset_for_linear_index(dims, strides, block, lhs);
229        for rhs in (lhs + 1)..total {
230            if block_offset_for_linear_index(dims, strides, block, rhs) == lhs_offset {
231                return false;
232            }
233        }
234    }
235    true
236}
237
238fn validate_injective_layout_inputs(dims: &[usize], strides: &[isize]) -> Option<usize> {
239    if dims.len() != strides.len() {
240        return None;
241    }
242
243    // A zero-sized layout is trivially injective even when the product of
244    // its other extents would overflow.
245    let total = crate::kernel::total_len(dims).ok()?;
246    if total <= 1 {
247        return Some(total);
248    }
249    if dims
250        .iter()
251        .zip(strides.iter())
252        .any(|(&dim, &stride)| dim > 1 && stride == 0)
253    {
254        return None;
255    }
256
257    let mut min_offset = 0isize;
258    let mut max_offset = 0isize;
259    for (&dim, &stride) in dims.iter().zip(strides.iter()) {
260        if dim <= 1 {
261            continue;
262        }
263        let extent = isize::try_from(dim - 1).ok()?;
264        let span = stride.checked_mul(extent)?;
265        if span >= 0 {
266            max_offset = max_offset.checked_add(span)?;
267        } else {
268            min_offset = min_offset.checked_add(span)?;
269        }
270    }
271    Some(total)
272}