strided_basic/
layout_check.rs1pub 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
42pub(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
66pub(crate) const EXACT_BLOCK_BUDGET: usize = 1 << 24;
74
75pub(crate) const PAIRWISE_BLOCK_BUDGET: usize = 4096;
81
82struct InterleavedBlock {
87 last_key: (u128, usize),
88 total: u128,
90 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 fn may_be_injective(&self) -> bool {
102 self.total <= self.span + 1
103 }
104}
105
106fn 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
150fn 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 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 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}