ruprim_host/expand/
mod.rs1use alloc::vec;
4use alloc::vec::Vec;
5use ruda_core::tensor::Shape;
6
7use ruda_core::tensor::host::{HostTensor, Layout};
8
9pub fn broadcast_shape(lhs: &Shape, rhs: &Shape) -> Shape {
13 let max_dims = lhs.num_dims().max(rhs.num_dims());
14 let mut result = vec![0; max_dims];
15
16 for (i, out) in result.iter_mut().enumerate() {
17 let lhs_idx = i as isize + lhs.num_dims() as isize - max_dims as isize;
18 let rhs_idx = i as isize + rhs.num_dims() as isize - max_dims as isize;
19
20 let lhs_dim = if lhs_idx >= 0 {
21 lhs[lhs_idx as usize]
22 } else {
23 1
24 };
25 let rhs_dim = if rhs_idx >= 0 {
26 rhs[rhs_idx as usize]
27 } else {
28 1
29 };
30
31 if lhs_dim == rhs_dim {
32 *out = lhs_dim;
33 } else if lhs_dim == 1 {
34 *out = rhs_dim;
35 } else if rhs_dim == 1 {
36 *out = lhs_dim;
37 } else {
38 panic!(
39 "broadcast_shape: incompatible dimensions {} and {} at position {}",
40 lhs_dim, rhs_dim, i
41 );
42 }
43 }
44
45 Shape::from(result)
46}
47
48pub fn broadcast_binary(lhs: HostTensor, rhs: HostTensor) -> (HostTensor, HostTensor) {
50 let lhs_shape = lhs.layout().shape().clone();
51 let rhs_shape = rhs.layout().shape().clone();
52
53 if lhs_shape == rhs_shape {
54 return (lhs, rhs);
55 }
56
57 let target = broadcast_shape(&lhs_shape, &rhs_shape);
58
59 let lhs_expanded = if lhs_shape == target {
60 lhs
61 } else {
62 expand(lhs, target.clone())
63 };
64 let rhs_expanded = if rhs_shape == target {
65 rhs
66 } else {
67 expand(rhs, target)
68 };
69
70 (lhs_expanded, rhs_expanded)
71}
72
73pub fn expand(tensor: HostTensor, target_shape: Shape) -> HostTensor {
78 let src_dims = tensor.layout().shape().to_vec();
80 let src_strides = tensor.layout().strides().to_vec();
81 let start_offset = tensor.layout().start_offset();
82 let dtype = tensor.dtype();
83
84 let src_ndims = src_dims.len();
85 let target_ndims = target_shape.num_dims();
86
87 assert!(
91 target_ndims >= src_ndims,
92 "expand: target rank ({}) must be >= source rank ({}); \
93 broadcasting cannot drop dimensions",
94 target_ndims,
95 src_ndims
96 );
97
98 let dim_diff = target_ndims - src_ndims;
100
101 let mut new_strides = Vec::with_capacity(target_ndims);
102
103 for i in 0..target_ndims {
104 let target_dim = target_shape[i];
105
106 if i < dim_diff {
107 new_strides.push(0);
109 } else {
110 let src_idx = i - dim_diff;
111 let src_dim = src_dims[src_idx];
112 let src_stride = src_strides[src_idx];
113
114 if src_dim == target_dim {
115 new_strides.push(src_stride);
117 } else if src_dim == 1 {
118 new_strides.push(0);
120 } else {
121 panic!(
122 "expand: cannot expand dimension {} from {} to {}",
123 i, src_dim, target_dim
124 );
125 }
126 }
127 }
128
129 let new_layout = Layout::new(target_shape, new_strides, start_offset);
130 HostTensor::from_arc(tensor.data_arc(), new_layout, dtype)
131}
132
133#[cfg(test)]
141mod tests;