Skip to main content

strided_perm/
fuse.rs

1//! Dimension fusion logic ported from Strided.jl/src/mapreduce.jl
2//!
3//! This module implements the core dimension fusion algorithm that merges
4//! contiguous dimensions to reduce iteration complexity.
5
6/// Fuse contiguous dimensions across multiple arrays.
7///
8/// This function fuses subsequent dimensions that are contiguous in memory
9/// for all arrays. If `strides[k][i] == dims[i-1] * strides[k][i-1]` for all k,
10/// dimensions i-1 and i can be merged.
11#[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    // Work from the end towards the beginning (Julia: for i in length(dims):-1:2)
21    for i in (1..n).rev() {
22        let mut can_merge = true;
23
24        // Check all arrays for contiguity
25        for strides in all_strides {
26            // s[i] should equal dims[i-1] * s[i-1] for fusion
27            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            // Fuse dimensions: merge dimension i into i-1
36            result[i - 1] *= result[i];
37            result[i] = 1;
38        }
39    }
40
41    result
42}
43
44/// Remove size-1 dimensions from fused dims and all corresponding strides.
45///
46/// After `fuse_dims()`, many dimensions may be 1 (either originally size-1
47/// or merged into a neighbor). These contribute nothing to iteration but
48/// increase loop depth. This function strips them out.
49///
50/// If ALL dimensions are 1 (scalar-like), a single dimension of size 1
51/// is preserved so the kernel has something to iterate over.
52#[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        // All dims are 1 (or empty). Preserve a single trivial dimension.
58        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/// Compute the "importance" of each dimension for loop ordering.
75///
76/// This encodes stride order information into importance scores that determine
77/// the optimal iteration order. The output array's strides are weighted 2x.
78#[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    // g = ceil(log2(M + 2)) = number of bits needed to encode array count
92    let g = (64 - (m as u64 + 1).leading_zeros()) as u64;
93
94    let mut importance = vec![0u64; n];
95
96    // First array (output) is weighted 2x
97    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    // Add contributions from remaining arrays
103    #[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    // Zero importance for size-1 dimensions (put them at the back)
112    for i in 0..n {
113        if dims[i] <= 1 {
114            importance[i] = 0;
115        }
116    }
117
118    importance
119}
120
121/// Get the permutation that sorts by importance (descending).
122#[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/// Compute the minimum stride cost for each dimension.
130#[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    // Transform: zero -> 1, nonzero -> 2*abs
147    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/// Validated result of fusing adjacent dimensions that are contiguous in
159/// both source and destination layouts.
160#[derive(Clone, Debug, PartialEq, Eq)]
161pub struct BilateralFusionPlan {
162    /// Fused non-trivial dimensions in iteration order.
163    pub dims: Vec<usize>,
164    /// Source strides corresponding to [`Self::dims`].
165    pub src_strides: Vec<isize>,
166    /// Destination strides corresponding to [`Self::dims`].
167    pub dst_strides: Vec<isize>,
168}
169
170/// Failure while validating or constructing a bilateral fusion plan.
171#[derive(Clone, Debug, thiserror::Error, PartialEq, Eq)]
172pub enum FusionPlanError {
173    /// Shape and stride ranks do not agree.
174    #[error(
175        "metadata lengths differ: dims={dims}, src_strides={src_strides}, \
176         dst_strides={dst_strides}"
177    )]
178    LengthMismatch {
179        /// Number of dimensions.
180        dims: usize,
181        /// Number of source strides.
182        src_strides: usize,
183        /// Number of destination strides.
184        dst_strides: usize,
185    },
186    /// A dimension cannot be represented for stride arithmetic, or a fused
187    /// dimension product exceeds `usize`.
188    #[error("fused dimension product overflows usize")]
189    DimensionOverflow,
190}
191
192/// Plan bilateral dimension fusion for source and destination layouts.
193///
194/// Size-one axes are removed. Two remaining adjacent axes are fused only when
195/// each layout is affine-contiguous across the boundary. Negative strides are
196/// supported; stride multiplication overflow simply prevents fusion.
197///
198/// # Errors
199///
200/// Returns [`FusionPlanError::LengthMismatch`] when the metadata ranks differ,
201/// and [`FusionPlanError::DimensionOverflow`] when a dimension cannot
202/// participate safely in `isize` stride arithmetic or a fused extent
203/// overflows `usize`.
204pub 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        // 24-dim all-size-2 col-major: all contiguous -> fuses to single dim
291        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        // src: contiguous 0-1, not 1-2
303        // dst: contiguous 0-1-2
304        let dims = vec![2, 3, 4];
305        let src_strides = vec![1, 2, 100]; // 0-1 contiguous, 1-2 not
306        let dst_strides = vec![1, 2, 6]; // all contiguous
307        let (fd, fs, fds) = fuse_dims_bilateral(&dims, &src_strides, &dst_strides);
308        assert_eq!(fd, vec![6, 4]); // first two fuse
309        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        // The benchmark case: scattered strides, nothing fuses
316        let dims = vec![2, 2, 2];
317        let src_strides = vec![1, 4194304, 2]; // scattered
318        let dst_strides = vec![1, 2, 4]; // contiguous
319        let (fd, fs, fds) = fuse_dims_bilateral(&dims, &src_strides, &dst_strides);
320        assert_eq!(fd, vec![2, 2, 2]); // nothing fuses
321        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}