strided_basic/
layout_check.rs1pub 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}