Skip to main content

strided_basic/
layout_check.rs

1//! Output layout validation shared by kernel families.
2
3/// Conservatively establish non-overlapping logical output positions.
4///
5/// # Examples
6///
7/// ```
8/// use strided_basic::execution::is_injective_layout;
9/// assert!(is_injective_layout(&[2, 3], &[1, 2]));
10/// assert!(!is_injective_layout(&[2, 3], &[0, 1]));
11/// ```
12pub fn is_injective_layout(dims: &[usize], strides: &[isize]) -> bool {
13    let Some(total) = validate_injective_layout_inputs(dims, strides) else {
14        return false;
15    };
16    if total <= 1 || has_disjoint_stride_spans(dims, strides) {
17        return true;
18    }
19
20    const EXACT_CHECK_LIMIT: usize = 4096;
21    if total <= EXACT_CHECK_LIMIT {
22        return has_unique_offsets_exact(dims, strides, total);
23    }
24
25    false
26}
27
28pub(crate) fn is_injective_layout_without_alloc(dims: &[usize], strides: &[isize]) -> bool {
29    let Some(total) = validate_injective_layout_inputs(dims, strides) else {
30        return false;
31    };
32    if total <= 1 || has_disjoint_stride_spans(dims, strides) {
33        return true;
34    }
35
36    const EXACT_CHECK_LIMIT: usize = 4096;
37    total <= EXACT_CHECK_LIMIT && has_unique_offsets_pairwise(dims, strides, total)
38}
39
40fn offset_for_linear_index(dims: &[usize], strides: &[isize], mut linear: usize) -> Option<isize> {
41    let mut offset = 0isize;
42    for (&dim, &stride) in dims.iter().zip(strides.iter()) {
43        let index = linear % dim;
44        linear /= dim;
45        offset = offset.checked_add(stride.checked_mul(index as isize)?)?;
46    }
47    Some(offset)
48}
49
50fn has_unique_offsets_pairwise(dims: &[usize], strides: &[isize], total: usize) -> bool {
51    for lhs in 0..total {
52        let Some(lhs_offset) = offset_for_linear_index(dims, strides, lhs) else {
53            return false;
54        };
55        for rhs in (lhs + 1)..total {
56            if offset_for_linear_index(dims, strides, rhs) == Some(lhs_offset) {
57                return false;
58            }
59        }
60    }
61    true
62}
63
64fn validate_injective_layout_inputs(dims: &[usize], strides: &[isize]) -> Option<usize> {
65    if dims.len() != strides.len() {
66        return None;
67    }
68
69    let total = dims
70        .iter()
71        .try_fold(1usize, |acc, &dim| acc.checked_mul(dim))?;
72    if total <= 1 {
73        return Some(total);
74    }
75    if dims
76        .iter()
77        .zip(strides.iter())
78        .any(|(&dim, &stride)| dim > 1 && stride == 0)
79    {
80        return None;
81    }
82
83    let mut min_offset = 0isize;
84    let mut max_offset = 0isize;
85    for (&dim, &stride) in dims.iter().zip(strides.iter()) {
86        if dim <= 1 {
87            continue;
88        }
89        let extent = isize::try_from(dim - 1).ok()?;
90        let span = stride.checked_mul(extent)?;
91        if span >= 0 {
92            max_offset = max_offset.checked_add(span)?;
93        } else {
94            min_offset = min_offset.checked_add(span)?;
95        }
96    }
97    Some(total)
98}
99
100fn has_unique_offsets_exact(dims: &[usize], strides: &[isize], total: usize) -> bool {
101    let mut seen = std::collections::HashSet::with_capacity(total);
102    let mut indices = vec![0usize; dims.len()];
103    let mut offset = 0isize;
104
105    for _ in 0..total {
106        if !seen.insert(offset) {
107            return false;
108        }
109
110        for axis in 0..dims.len() {
111            indices[axis] += 1;
112            offset = match offset.checked_add(strides[axis]) {
113                Some(offset) => offset,
114                None => return false,
115            };
116            if indices[axis] < dims[axis] {
117                break;
118            }
119
120            let rewind = match strides[axis].checked_mul(indices[axis] as isize) {
121                Some(rewind) => rewind,
122                None => return false,
123            };
124            offset = match offset.checked_sub(rewind) {
125                Some(offset) => offset,
126                None => return false,
127            };
128            indices[axis] = 0;
129        }
130    }
131
132    true
133}
134
135fn has_disjoint_stride_spans(dims: &[usize], strides: &[isize]) -> bool {
136    let mut covered_span = 0u128;
137    let mut previous_axis = None;
138    let active_axes = dims.iter().filter(|&&dim| dim > 1).count();
139    for _ in 0..active_axes {
140        let mut next = None;
141        for (axis, (&dim, &stride)) in dims.iter().zip(strides.iter()).enumerate() {
142            if dim <= 1 {
143                continue;
144            }
145            let stride = match stride.checked_abs() {
146                Some(stride) => stride as u128,
147                None => return false,
148            };
149            let key = (stride, axis);
150            if previous_axis.is_some_and(|previous| key <= previous) {
151                continue;
152            }
153            if next.is_none_or(|(best, _)| key < best) {
154                next = Some((key, dim as u128 - 1));
155            }
156        }
157        let Some(((stride, axis), extent)) = next else {
158            return false;
159        };
160        if stride <= covered_span {
161            return false;
162        }
163        covered_span = match stride
164            .checked_mul(extent)
165            .and_then(|span| covered_span.checked_add(span))
166        {
167            Some(covered_span) => covered_span,
168            None => return false,
169        };
170        previous_axis = Some((stride, axis));
171    }
172
173    true
174}