use alloc::vec;
use alloc::vec::Vec;
use ruda_core::tensor::Shape;
use ruda_core::tensor::host::{HostTensor, Layout};
pub fn broadcast_shape(lhs: &Shape, rhs: &Shape) -> Shape {
let max_dims = lhs.num_dims().max(rhs.num_dims());
let mut result = vec![0; max_dims];
for (i, out) in result.iter_mut().enumerate() {
let lhs_idx = i as isize + lhs.num_dims() as isize - max_dims as isize;
let rhs_idx = i as isize + rhs.num_dims() as isize - max_dims as isize;
let lhs_dim = if lhs_idx >= 0 {
lhs[lhs_idx as usize]
} else {
1
};
let rhs_dim = if rhs_idx >= 0 {
rhs[rhs_idx as usize]
} else {
1
};
if lhs_dim == rhs_dim {
*out = lhs_dim;
} else if lhs_dim == 1 {
*out = rhs_dim;
} else if rhs_dim == 1 {
*out = lhs_dim;
} else {
panic!(
"broadcast_shape: incompatible dimensions {} and {} at position {}",
lhs_dim, rhs_dim, i
);
}
}
Shape::from(result)
}
pub fn broadcast_binary(lhs: HostTensor, rhs: HostTensor) -> (HostTensor, HostTensor) {
let lhs_shape = lhs.layout().shape().clone();
let rhs_shape = rhs.layout().shape().clone();
if lhs_shape == rhs_shape {
return (lhs, rhs);
}
let target = broadcast_shape(&lhs_shape, &rhs_shape);
let lhs_expanded = if lhs_shape == target {
lhs
} else {
expand(lhs, target.clone())
};
let rhs_expanded = if rhs_shape == target {
rhs
} else {
expand(rhs, target)
};
(lhs_expanded, rhs_expanded)
}
pub fn expand(tensor: HostTensor, target_shape: Shape) -> HostTensor {
let src_dims = tensor.layout().shape().to_vec();
let src_strides = tensor.layout().strides().to_vec();
let start_offset = tensor.layout().start_offset();
let dtype = tensor.dtype();
let src_ndims = src_dims.len();
let target_ndims = target_shape.num_dims();
assert!(
target_ndims >= src_ndims,
"expand: target rank ({}) must be >= source rank ({}); \
broadcasting cannot drop dimensions",
target_ndims,
src_ndims
);
let dim_diff = target_ndims - src_ndims;
let mut new_strides = Vec::with_capacity(target_ndims);
for i in 0..target_ndims {
let target_dim = target_shape[i];
if i < dim_diff {
new_strides.push(0);
} else {
let src_idx = i - dim_diff;
let src_dim = src_dims[src_idx];
let src_stride = src_strides[src_idx];
if src_dim == target_dim {
new_strides.push(src_stride);
} else if src_dim == 1 {
new_strides.push(0);
} else {
panic!(
"expand: cannot expand dimension {} from {} to {}",
i, src_dim, target_dim
);
}
}
}
let new_layout = Layout::new(target_shape, new_strides, start_offset);
HostTensor::from_arc(tensor.data_arc(), new_layout, dtype)
}
#[cfg(test)]
mod tests;