Skip to main content

ruprim_host/expand/
mod.rs

1//! Expand operation for broadcasting tensors to larger shapes.
2
3use alloc::vec;
4use alloc::vec::Vec;
5use ruda_core::tensor::Shape;
6
7use ruda_core::tensor::host::{HostTensor, Layout};
8
9/// Compute the broadcast shape of two tensors.
10///
11/// Returns the shape that both tensors can be expanded to for element-wise operations.
12pub 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
48/// Broadcast two tensors to the same shape for binary operations.
49pub 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
73/// Expand a tensor to a larger shape by broadcasting.
74///
75/// Dimensions of size 1 can be expanded to any size. The result is a view
76/// that doesn't copy data - it uses stride 0 for expanded dimensions.
77pub fn expand(tensor: HostTensor, target_shape: Shape) -> HostTensor {
78    // Capture values we need before consuming tensor
79    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    // Broadcasting only prepends new dims; it never drops existing ones.
88    // A target with fewer dims than the source would silently discard the
89    // trailing source strides and produce an invalid layout.
90    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    // Prepend 1s to source shape if needed (for broadcasting like [3] -> [2, 3])
99    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 dimension prepended - must be broadcastable from size 1
108            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                // Same size - keep stride
116                new_strides.push(src_stride);
117            } else if src_dim == 1 {
118                // Broadcast dimension - stride becomes 0
119                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// Tests kept here probe flex-internal expand behavior: stride metadata
134// (stride 0 on broadcast dims, preservation of negative strides on
135// flipped inputs, preserved start-offset on narrowed inputs) and the
136// flex-only `broadcast_binary` helper. Public-API expand coverage for
137// transpose/flip/narrow variants lives in
138// crates/ruda-backend-tests/tests/tensor/{float,int,bool}/ops/expand.rs
139// so it runs on every backend.
140#[cfg(test)]
141mod tests;