1use alloc::vec::Vec;
17
18use crate::Shape;
19
20pub type DimOrder = Shape;
25
26pub fn dim_order(shape: &[usize], strides: &[usize]) -> Option<DimOrder> {
37 dim_order_inner(shape, strides, Padding::Rejected)
38}
39
40pub fn nested_dim_order(shape: &[usize], strides: &[usize]) -> Option<DimOrder> {
56 dim_order_inner(shape, strides, Padding::Allowed)
57}
58
59#[derive(Clone, Copy)]
60enum Padding {
61 Allowed,
62 Rejected,
63}
64
65fn dim_order_inner(shape: &[usize], strides: &[usize], padding: Padding) -> Option<DimOrder> {
66 let rank = shape.len();
67
68 if rank != strides.len() {
69 return None;
70 }
71
72 let mut order: Vec<usize> = (0..rank).collect();
73 order.sort_by(|a, b| strides[*b].cmp(&strides[*a]).then(a.cmp(b)));
77
78 let mut expected = 1;
79
80 for &axis in order.iter().rev() {
81 if shape[axis] == 1 {
82 continue;
83 }
84 match padding {
85 Padding::Allowed if strides[axis] < expected => return None,
88 Padding::Rejected if strides[axis] != expected => return None,
89 _ => {}
90 }
91 expected = strides[axis] * shape[axis];
94 }
95
96 Some(Shape::from(order))
97}
98
99pub fn is_contiguous_order(order: &[usize]) -> bool {
101 order.iter().enumerate().all(|(pos, axis)| pos == *axis)
102}
103
104#[cfg(test)]
105mod tests {
106 use super::*;
107 use alloc::vec;
108
109 #[test]
110 fn contiguous_is_the_identity_order() {
111 let shape = [2, 48, 16, 16];
112 let strides = [48 * 16 * 16, 16 * 16, 16, 1];
113
114 assert_eq!(
115 dim_order(&shape, &strides),
116 Some(Shape::from(vec![0, 1, 2, 3]))
117 );
118 }
119
120 #[test]
121 fn nhwc_memory_gives_the_nhwc_order() {
122 let shape = [2, 48, 16, 16];
124 let strides = [16 * 16 * 48, 1, 16 * 48, 48];
125
126 assert_eq!(
127 dim_order(&shape, &strides),
128 Some(Shape::from(vec![0, 2, 3, 1]))
129 );
130 }
131
132 #[test]
133 fn broadcast_is_not_dense() {
134 let shape = [2, 48, 16, 16];
135 let strides = [0, 1, 0, 0];
136
137 assert_eq!(dim_order(&shape, &strides), None);
138 }
139
140 #[test]
141 fn a_slice_is_not_dense() {
142 let shape = [4, 8];
144 let strides = [16, 1];
145
146 assert_eq!(dim_order(&shape, &strides), None);
147 }
148
149 #[test]
150 fn size_one_dimensions_do_not_decide_density() {
151 let shape = [1, 48, 1, 1];
154 let strides = [48, 1, 48, 48];
155
156 assert!(dim_order(&shape, &strides).is_some());
157 }
158
159 #[test]
160 fn the_order_ends_at_the_innermost_dimension() {
161 let shape = [2, 48, 16, 16];
162
163 let contiguous = dim_order(&shape, &[48 * 16 * 16, 16 * 16, 16, 1]).unwrap();
164 let nhwc = dim_order(&shape, &[16 * 16 * 48, 1, 16 * 48, 48]).unwrap();
165
166 assert_eq!(contiguous.last(), Some(&3));
167 assert_eq!(nhwc.last(), Some(&1));
168 }
169
170 #[test]
171 fn order_is_a_permutation() {
172 assert!(is_contiguous_order(&[0, 1, 2, 3]));
173 assert!(!is_contiguous_order(&[0, 2, 3, 1]));
174 }
175
176 #[test]
177 fn a_dense_tensor_nests_in_the_order_it_is_dense_in() {
178 let shape = [2, 48, 16, 16];
179
180 for strides in [
181 [48 * 16 * 16, 16 * 16, 16, 1],
182 [16 * 16 * 48, 1, 16 * 48, 48],
183 ] {
184 let dense = dim_order(&shape, &strides);
185 assert!(dense.is_some());
186 assert_eq!(nested_dim_order(&shape, &strides), dense);
187 }
188 }
189
190 #[test]
191 fn padding_under_the_innermost_dimension_keeps_the_nhwc_order() {
192 let shape = [2, 48, 16, 16];
193 let strides = [16 * 16 * 64, 1, 16 * 64, 64];
194
195 assert_eq!(dim_order(&shape, &strides), None);
196 assert_eq!(
197 nested_dim_order(&shape, &strides),
198 Some(Shape::from(vec![0, 2, 3, 1]))
199 );
200 }
201
202 #[test]
203 fn padding_above_the_innermost_dimension_is_nesting_too() {
204 let shape = [4, 8];
205 let strides = [16, 1];
206
207 assert_eq!(dim_order(&shape, &strides), None);
208 assert_eq!(
209 nested_dim_order(&shape, &strides),
210 Some(Shape::from(vec![0, 1]))
211 );
212 }
213
214 #[test]
215 fn overlapping_dimensions_are_not_an_order_under_either() {
216 let shape = [4, 8];
217 let strides = [4, 1];
218
219 assert_eq!(dim_order(&shape, &strides), None);
220 assert_eq!(nested_dim_order(&shape, &strides), None);
221 }
222
223 #[test]
224 fn a_broadcast_dimension_still_cannot_vote() {
225 let shape = [2, 48, 16, 16];
226 let strides = [0, 1, 0, 0];
227
228 assert_eq!(nested_dim_order(&shape, &strides), None);
229 }
230
231 #[test]
232 fn size_one_dimensions_neither_pad_nor_constrain_nesting() {
233 let shape = [1, 48, 1, 1];
234 let strides = [48, 1, 48, 48];
235
236 assert_eq!(
237 nested_dim_order(&shape, &strides),
238 dim_order(&shape, &strides)
239 );
240 assert!(nested_dim_order(&shape, &strides).is_some());
241 }
242}