1pub 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
38pub 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
49pub 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
61pub 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#[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 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 pub fn len(&self) -> usize {
166 self.len
167 }
168
169 pub fn is_empty(&self) -> bool {
171 self.len == 0
172 }
173
174 pub fn output_shape(&self) -> &[usize] {
176 &self.output_shape
177 }
178
179 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
191pub 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}