1#[cfg(test)]
12pub fn fuse_dims(dims: &[usize], all_strides: &[&[isize]]) -> Vec<usize> {
13 let n = dims.len();
14 if n <= 1 || all_strides.is_empty() {
15 return dims.to_vec();
16 }
17
18 let mut result = dims.to_vec();
19
20 for i in (1..n).rev() {
22 let mut can_merge = true;
23
24 for strides in all_strides {
26 let expected = result[i - 1] as isize * strides[i - 1];
28 if strides[i] != expected {
29 can_merge = false;
30 break;
31 }
32 }
33
34 if can_merge {
35 result[i - 1] *= result[i];
37 result[i] = 1;
38 }
39 }
40
41 result
42}
43
44#[cfg(test)]
53pub fn compress_dims(dims: &[usize], all_strides: &[Vec<isize>]) -> (Vec<usize>, Vec<Vec<isize>>) {
54 let kept: Vec<usize> = (0..dims.len()).filter(|&i| dims[i] != 1).collect();
55
56 if kept.is_empty() {
57 if dims.is_empty() {
59 return (vec![], all_strides.to_vec());
60 }
61 let new_strides = all_strides.iter().map(|s| vec![s[0]]).collect();
62 return (vec![1], new_strides);
63 }
64
65 let new_dims: Vec<usize> = kept.iter().map(|&i| dims[i]).collect();
66 let new_strides: Vec<Vec<isize>> = all_strides
67 .iter()
68 .map(|s| kept.iter().map(|&i| s[i]).collect())
69 .collect();
70
71 (new_dims, new_strides)
72}
73
74#[cfg(test)]
79pub fn compute_importance(
80 dims: &[usize],
81 all_strides: &[&[isize]],
82 index_orders: &[Vec<usize>],
83) -> Vec<u64> {
84 let n = dims.len();
85 let m = all_strides.len();
86
87 if n == 0 || m == 0 {
88 return vec![];
89 }
90
91 let g = (64 - (m as u64 + 1).leading_zeros()) as u64;
93
94 let mut importance = vec![0u64; n];
95
96 for i in 0..n {
98 let shift = g * (n - index_orders[0][i]) as u64;
99 importance[i] = 2 * (1u64 << shift);
100 }
101
102 #[allow(clippy::needless_range_loop)]
104 for k in 1..m {
105 for i in 0..n {
106 let shift = g * (n - index_orders[k][i]) as u64;
107 importance[i] += 1u64 << shift;
108 }
109 }
110
111 for i in 0..n {
113 if dims[i] <= 1 {
114 importance[i] = 0;
115 }
116 }
117
118 importance
119}
120
121#[cfg(test)]
123pub fn sort_by_importance(importance: &[u64]) -> Vec<usize> {
124 let mut indices: Vec<usize> = (0..importance.len()).collect();
125 indices.sort_by(|&a, &b| importance[b].cmp(&importance[a]));
126 indices
127}
128
129#[cfg(test)]
131pub fn compute_costs<S: AsRef<[isize]>>(all_strides: &[S]) -> Vec<isize> {
132 if all_strides.is_empty() {
133 return vec![];
134 }
135
136 let n = all_strides[0].as_ref().len();
137 let mut costs = vec![isize::MAX; n];
138
139 for strides in all_strides {
140 let strides = strides.as_ref();
141 for i in 0..n {
142 costs[i] = costs[i].min(strides[i].abs());
143 }
144 }
145
146 for cost in &mut costs {
148 if *cost == 0 {
149 *cost = 1;
150 } else {
151 *cost *= 2;
152 }
153 }
154
155 costs
156}
157
158#[derive(Clone, Debug, PartialEq, Eq)]
161pub struct BilateralFusionPlan {
162 pub dims: Vec<usize>,
164 pub src_strides: Vec<isize>,
166 pub dst_strides: Vec<isize>,
168}
169
170#[derive(Clone, Debug, thiserror::Error, PartialEq, Eq)]
172pub enum FusionPlanError {
173 #[error(
175 "metadata lengths differ: dims={dims}, src_strides={src_strides}, \
176 dst_strides={dst_strides}"
177 )]
178 LengthMismatch {
179 dims: usize,
181 src_strides: usize,
183 dst_strides: usize,
185 },
186 #[error("fused dimension product overflows usize")]
189 DimensionOverflow,
190}
191
192pub fn plan_bilateral_fusion(
205 dims: &[usize],
206 src_strides: &[isize],
207 dst_strides: &[isize],
208) -> Result<BilateralFusionPlan, FusionPlanError> {
209 if dims.len() != src_strides.len() || dims.len() != dst_strides.len() {
210 return Err(FusionPlanError::LengthMismatch {
211 dims: dims.len(),
212 src_strides: src_strides.len(),
213 dst_strides: dst_strides.len(),
214 });
215 }
216
217 let mut plan = BilateralFusionPlan {
218 dims: Vec::with_capacity(dims.len()),
219 src_strides: Vec::with_capacity(dims.len()),
220 dst_strides: Vec::with_capacity(dims.len()),
221 };
222
223 for ((&dim, &src_stride), &dst_stride) in dims.iter().zip(src_strides).zip(dst_strides) {
224 if dim == 1 {
225 continue;
226 }
227 isize::try_from(dim).map_err(|_| FusionPlanError::DimensionOverflow)?;
228
229 if let Some(last) = plan.dims.len().checked_sub(1) {
230 let current_dim =
231 isize::try_from(plan.dims[last]).map_err(|_| FusionPlanError::DimensionOverflow)?;
232 let src_contiguous = plan.src_strides[last]
233 .checked_mul(current_dim)
234 .is_some_and(|expected| src_stride == expected);
235 let dst_contiguous = plan.dst_strides[last]
236 .checked_mul(current_dim)
237 .is_some_and(|expected| dst_stride == expected);
238
239 if src_contiguous && dst_contiguous {
240 plan.dims[last] = plan.dims[last]
241 .checked_mul(dim)
242 .ok_or(FusionPlanError::DimensionOverflow)?;
243 continue;
244 }
245 }
246
247 plan.dims.push(dim);
248 plan.src_strides.push(src_stride);
249 plan.dst_strides.push(dst_stride);
250 }
251
252 Ok(plan)
253}
254
255#[cfg(test)]
256fn fuse_dims_bilateral(
257 dims: &[usize],
258 src_strides: &[isize],
259 dst_strides: &[isize],
260) -> (Vec<usize>, Vec<isize>, Vec<isize>) {
261 let plan = plan_bilateral_fusion(dims, src_strides, dst_strides).unwrap();
262 (plan.dims, plan.src_strides, plan.dst_strides)
263}
264
265#[cfg(test)]
266mod tests {
267 use super::*;
268
269 #[test]
270 fn test_fuse_dims_contiguous() {
271 let dims = [3, 4];
272 let strides1 = [1isize, 3];
273 let strides2 = [1isize, 3];
274 let all_strides: Vec<&[isize]> = vec![&strides1, &strides2];
275 let fused = fuse_dims(&dims, &all_strides);
276 assert_eq!(fused, vec![12, 1]);
277 }
278
279 #[test]
280 fn test_fuse_dims_non_contiguous() {
281 let dims = [3, 4];
282 let strides1 = [1isize, 10];
283 let all_strides: Vec<&[isize]> = vec![&strides1];
284 let fused = fuse_dims(&dims, &all_strides);
285 assert_eq!(fused, vec![3, 4]);
286 }
287
288 #[test]
289 fn test_fuse_dims_bilateral_all_contiguous() {
290 let dims = vec![2, 2, 2, 2];
292 let src_strides = vec![1, 2, 4, 8];
293 let dst_strides = vec![1, 2, 4, 8];
294 let (fd, fs, fds) = fuse_dims_bilateral(&dims, &src_strides, &dst_strides);
295 assert_eq!(fd, vec![16]);
296 assert_eq!(fs, vec![1]);
297 assert_eq!(fds, vec![1]);
298 }
299
300 #[test]
301 fn test_fuse_dims_bilateral_partial() {
302 let dims = vec![2, 3, 4];
305 let src_strides = vec![1, 2, 100]; let dst_strides = vec![1, 2, 6]; let (fd, fs, fds) = fuse_dims_bilateral(&dims, &src_strides, &dst_strides);
308 assert_eq!(fd, vec![6, 4]); assert_eq!(fs, vec![1, 100]);
310 assert_eq!(fds, vec![1, 6]);
311 }
312
313 #[test]
314 fn test_fuse_dims_bilateral_scattered() {
315 let dims = vec![2, 2, 2];
317 let src_strides = vec![1, 4194304, 2]; let dst_strides = vec![1, 2, 4]; let (fd, fs, fds) = fuse_dims_bilateral(&dims, &src_strides, &dst_strides);
320 assert_eq!(fd, vec![2, 2, 2]); assert_eq!(fs, vec![1, 4194304, 2]);
322 assert_eq!(fds, vec![1, 2, 4]);
323 }
324
325 #[test]
326 fn test_compress_dims_removes_fused() {
327 let dims = vec![12usize, 1];
328 let strides = vec![vec![1isize, 3]];
329 let (cd, cs) = compress_dims(&dims, &strides);
330 assert_eq!(cd, vec![12]);
331 assert_eq!(cs, vec![vec![1]]);
332 }
333
334 #[test]
335 fn test_compute_costs() {
336 let strides1 = [1isize, 4, 0];
337 let strides2 = [2isize, 1, 0];
338 let all_strides: Vec<&[isize]> = vec![&strides1, &strides2];
339 let costs = compute_costs(&all_strides);
340 assert_eq!(costs, vec![2, 2, 1]);
341 }
342}