Skip to main content

runmat_runtime/builtins/common/
broadcast.rs

1//! Broadcasting utilities shared across builtin implementations.
2//!
3//! The helpers in this module mirror MATLAB's implicit expansion rules and
4//! operate on column-major shapes expressed as `[usize]` vectors.
5
6/// Compute the broadcasted shape for two operands using MATLAB implicit
7/// expansion rules.
8pub fn broadcast_shapes(
9    fn_name: &str,
10    left: &[usize],
11    right: &[usize],
12) -> Result<Vec<usize>, String> {
13    let rank = left.len().max(right.len());
14    let left = align_shape(left, rank);
15    let right = align_shape(right, rank);
16    let mut shape = Vec::with_capacity(rank);
17    for dim in 0..rank {
18        let a = left[dim];
19        let b = right[dim];
20        if a == b {
21            shape.push(a);
22        } else if a == 1 {
23            shape.push(b);
24        } else if b == 1 {
25            shape.push(a);
26        } else {
27            return Err(format!(
28                "{fn_name}: size mismatch between inputs (dimension {} has lengths {} and {})",
29                dim + 1,
30                a,
31                b
32            ));
33        }
34    }
35    Ok(shape)
36}
37
38/// Compute column-major strides for a given shape.
39pub fn compute_strides(shape: &[usize]) -> Vec<usize> {
40    let mut strides = Vec::with_capacity(shape.len());
41    let mut stride = 1usize;
42    for &extent in shape {
43        strides.push(stride);
44        stride = stride.saturating_mul(extent.max(1));
45    }
46    strides
47}
48
49/// Append MATLAB's implicit trailing singleton dimensions until `shape` reaches `rank`.
50pub fn align_shape(shape: &[usize], rank: usize) -> Vec<usize> {
51    debug_assert!(shape.len() <= rank);
52    let mut aligned = if shape.len() == 1 && rank >= 2 {
53        vec![1, shape[0]]
54    } else {
55        shape.to_vec()
56    };
57    aligned.resize(rank, 1);
58    aligned
59}
60
61/// Map a linear index in the broadcasted result back to a source operand.
62pub fn broadcast_index(
63    mut linear: usize,
64    out_shape: &[usize],
65    in_shape: &[usize],
66    strides: &[usize],
67) -> usize {
68    if in_shape.is_empty() {
69        return 0;
70    }
71    let row_shorthand = in_shape.len() == 1 && out_shape.len() >= 2;
72    let mut offset = 0usize;
73    for dim in 0..out_shape.len() {
74        let out_extent = out_shape[dim];
75        let coord = if out_extent == 0 {
76            0
77        } else {
78            linear % out_extent
79        };
80        if out_extent != 0 {
81            linear /= out_extent;
82        }
83        let (in_extent, in_stride) = if row_shorthand {
84            match dim {
85                1 => (in_shape[0], 1),
86                _ => (1, 0),
87            }
88        } else {
89            (
90                in_shape.get(dim).copied().unwrap_or(1),
91                strides.get(dim).copied().unwrap_or(0),
92            )
93        };
94        let mapped = if in_extent == 1 || out_extent == 0 {
95            0
96        } else {
97            coord
98        };
99        offset += mapped * in_stride;
100    }
101    offset
102}
103
104/// Broadcast plan describing how two tensors can be implicitly expanded.
105#[derive(Debug, Clone)]
106pub struct BroadcastPlan {
107    output_shape: Vec<usize>,
108    len: usize,
109    advance_a: Vec<usize>,
110    advance_b: Vec<usize>,
111}
112
113impl BroadcastPlan {
114    /// Construct a broadcast plan for two shapes, returning an error when they
115    /// cannot be implicitly expanded under MATLAB rules.
116    pub fn new(shape_a: &[usize], shape_b: &[usize]) -> Result<Self, String> {
117        let ndims = shape_a.len().max(shape_b.len());
118
119        let ext_a = align_shape(shape_a, ndims);
120        let ext_b = align_shape(shape_b, ndims);
121        let mut output_shape = Vec::with_capacity(ndims);
122        for i in 0..ndims {
123            let da = ext_a[i];
124            let db = ext_b[i];
125            if da == db {
126                output_shape.push(da);
127            } else if da == 1 {
128                output_shape.push(db);
129            } else if db == 1 {
130                output_shape.push(da);
131            } else {
132                return Err(format!(
133                    "broadcast: non-singleton dimension mismatch (dimension {}: {} vs {})",
134                    i + 1,
135                    da,
136                    db
137                ));
138            }
139        }
140
141        let len = output_shape.iter().copied().product();
142        let strides_a = compute_strides(&ext_a);
143        let strides_b = compute_strides(&ext_b);
144
145        let advance_a = ext_a
146            .iter()
147            .enumerate()
148            .map(|(dim, &size)| if size <= 1 { 0 } else { strides_a[dim] })
149            .collect::<Vec<_>>();
150        let advance_b = ext_b
151            .iter()
152            .enumerate()
153            .map(|(dim, &size)| if size <= 1 { 0 } else { strides_b[dim] })
154            .collect::<Vec<_>>();
155
156        Ok(Self {
157            output_shape,
158            len,
159            advance_a,
160            advance_b,
161        })
162    }
163
164    /// Total number of elements produced by the broadcast.
165    pub fn len(&self) -> usize {
166        self.len
167    }
168
169    /// Returns true if the broadcast produces no elements.
170    pub fn is_empty(&self) -> bool {
171        self.len == 0
172    }
173
174    /// Output shape after broadcasting both operands.
175    pub fn output_shape(&self) -> &[usize] {
176        &self.output_shape
177    }
178
179    /// Iterator yielding `(output_index, index_a, index_b)` triples for each element.
180    pub fn iter(&self) -> BroadcastIter<'_> {
181        BroadcastIter {
182            plan: self,
183            offset: 0,
184            index_a: 0,
185            index_b: 0,
186            coords: vec![0usize; self.output_shape.len()],
187        }
188    }
189}
190
191/// Iterator over broadcast indices.
192pub struct BroadcastIter<'a> {
193    plan: &'a BroadcastPlan,
194    offset: usize,
195    index_a: usize,
196    index_b: usize,
197    coords: Vec<usize>,
198}
199
200impl<'a> Iterator for BroadcastIter<'a> {
201    type Item = (usize, usize, usize);
202
203    fn next(&mut self) -> Option<Self::Item> {
204        if self.offset >= self.plan.len {
205            return None;
206        }
207        let current = (self.offset, self.index_a, self.index_b);
208        self.offset += 1;
209        if self.offset == self.plan.len {
210            return Some(current);
211        }
212        for dim in 0..self.plan.output_shape.len() {
213            if self.plan.output_shape[dim] == 0 {
214                continue;
215            }
216            self.coords[dim] += 1;
217            if self.coords[dim] < self.plan.output_shape[dim] {
218                self.index_a += self.plan.advance_a[dim];
219                self.index_b += self.plan.advance_b[dim];
220                break;
221            }
222            self.coords[dim] = 0;
223            let rewind = self.plan.output_shape[dim].saturating_sub(1);
224            let rewind_a = self.plan.advance_a[dim] * rewind;
225            let rewind_b = self.plan.advance_b[dim] * rewind;
226            if rewind_a != 0 {
227                self.index_a = self.index_a.saturating_sub(rewind_a);
228            }
229            if rewind_b != 0 {
230                self.index_b = self.index_b.saturating_sub(rewind_b);
231            }
232        }
233        Some(current)
234    }
235}
236
237#[cfg(test)]
238pub(crate) mod tests {
239    use super::*;
240
241    #[cfg_attr(target_arch = "wasm32", wasm_bindgen_test::wasm_bindgen_test)]
242    #[test]
243    fn broadcast_equal_shapes() {
244        let out = broadcast_shapes("test", &[2, 3], &[2, 3]).unwrap();
245        assert_eq!(out, vec![2, 3]);
246    }
247
248    #[cfg_attr(target_arch = "wasm32", wasm_bindgen_test::wasm_bindgen_test)]
249    #[test]
250    fn broadcast_scalar() {
251        let out = broadcast_shapes("test", &[1, 1], &[4, 5]).unwrap();
252        assert_eq!(out, vec![4, 5]);
253    }
254
255    #[cfg_attr(target_arch = "wasm32", wasm_bindgen_test::wasm_bindgen_test)]
256    #[test]
257    fn broadcast_mismatched_dimension_errors() {
258        let err = broadcast_shapes("test", &[2, 3], &[4, 3]).unwrap_err();
259        assert!(err.contains("dimension 1"));
260    }
261
262    #[test]
263    fn broadcast_appends_missing_trailing_singletons() {
264        assert_eq!(
265            broadcast_shapes("test", &[2, 3], &[2, 3, 4]).unwrap(),
266            vec![2, 3, 4]
267        );
268        let error = broadcast_shapes("test", &[2, 3], &[1, 2, 3]).unwrap_err();
269        assert!(error.contains("dimension 2"));
270    }
271
272    #[test]
273    fn broadcast_zero_is_compatible_only_with_zero_or_one() {
274        assert_eq!(
275            broadcast_shapes("test", &[0, 3], &[1, 3]).unwrap(),
276            vec![0, 3]
277        );
278        assert_eq!(
279            broadcast_shapes("test", &[0, 3], &[0, 3]).unwrap(),
280            vec![0, 3]
281        );
282        assert!(broadcast_shapes("test", &[0, 3], &[2, 3]).is_err());
283    }
284
285    #[cfg_attr(target_arch = "wasm32", wasm_bindgen_test::wasm_bindgen_test)]
286    #[test]
287    fn compute_strides_column_major() {
288        let strides = compute_strides(&[2, 3, 4]);
289        assert_eq!(strides, vec![1, 2, 6]);
290    }
291
292    #[test]
293    fn align_shape_appends_trailing_singletons() {
294        assert_eq!(align_shape(&[2, 3], 4), vec![2, 3, 1, 1]);
295        assert_eq!(align_shape(&[3], 3), vec![1, 3, 1]);
296        assert_eq!(align_shape(&[3], 1), vec![3]);
297    }
298
299    #[cfg_attr(target_arch = "wasm32", wasm_bindgen_test::wasm_bindgen_test)]
300    #[test]
301    fn broadcast_index_maps_scalar_inputs() {
302        let strides = compute_strides(&[1, 1]);
303        let idx = broadcast_index(5, &[2, 3], &[1, 1], &strides);
304        assert_eq!(idx, 0);
305    }
306
307    #[test]
308    fn broadcast_index_maps_one_dimensional_row_shorthand() {
309        let strides = compute_strides(&[3]);
310        assert_eq!(
311            (0..6)
312                .map(|linear| broadcast_index(linear, &[1, 3, 2], &[3], &strides))
313                .collect::<Vec<_>>(),
314            vec![0, 1, 2, 0, 1, 2]
315        );
316    }
317
318    #[cfg_attr(target_arch = "wasm32", wasm_bindgen_test::wasm_bindgen_test)]
319    #[test]
320    fn broadcast_same_shape() {
321        let plan = BroadcastPlan::new(&[2, 3], &[2, 3]).unwrap();
322        assert_eq!(plan.output_shape(), &[2, 3]);
323        assert_eq!(plan.len(), 6);
324        let indices: Vec<(usize, usize, usize)> = plan.iter().collect();
325        assert_eq!(
326            indices,
327            vec![
328                (0, 0, 0),
329                (1, 1, 1),
330                (2, 2, 2),
331                (3, 3, 3),
332                (4, 4, 4),
333                (5, 5, 5),
334            ]
335        );
336    }
337
338    #[cfg_attr(target_arch = "wasm32", wasm_bindgen_test::wasm_bindgen_test)]
339    #[test]
340    fn broadcast_scalar_expansion() {
341        let plan = BroadcastPlan::new(&[1, 3], &[1, 1]).unwrap();
342        assert_eq!(plan.output_shape(), &[1, 3]);
343        assert_eq!(plan.len(), 3);
344        let indices: Vec<(usize, usize, usize)> = plan.iter().collect();
345        assert_eq!(indices, vec![(0, 0, 0), (1, 1, 0), (2, 2, 0)]);
346    }
347
348    #[cfg_attr(target_arch = "wasm32", wasm_bindgen_test::wasm_bindgen_test)]
349    #[test]
350    fn broadcast_zero_sized_dimension() {
351        let plan = BroadcastPlan::new(&[0, 3], &[1, 3]).unwrap();
352        assert_eq!(plan.output_shape(), &[0, 3]);
353        assert_eq!(plan.len(), 0);
354        assert_eq!(plan.iter().next(), None);
355    }
356
357    #[test]
358    fn broadcast_plan_appends_missing_trailing_singletons() {
359        let plan = BroadcastPlan::new(&[2, 1], &[2, 1, 3]).unwrap();
360        assert_eq!(plan.output_shape(), &[2, 1, 3]);
361        assert_eq!(
362            plan.iter().collect::<Vec<_>>(),
363            vec![
364                (0, 0, 0),
365                (1, 1, 1),
366                (2, 0, 2),
367                (3, 1, 3),
368                (4, 0, 4),
369                (5, 1, 5),
370            ]
371        );
372        assert!(BroadcastPlan::new(&[2, 3], &[1, 2, 3]).is_err());
373        assert!(BroadcastPlan::new(&[0, 3], &[2, 3]).is_err());
374        let row_shorthand = BroadcastPlan::new(&[3], &[1, 3, 2]).unwrap();
375        assert_eq!(row_shorthand.output_shape(), &[1, 3, 2]);
376        assert_eq!(
377            row_shorthand.iter().collect::<Vec<_>>(),
378            vec![
379                (0, 0, 0),
380                (1, 1, 1),
381                (2, 2, 2),
382                (3, 0, 3),
383                (4, 1, 4),
384                (5, 2, 5),
385            ]
386        );
387    }
388}