use crate::dsl::Runtime;
use super::RudaTensor;
use ruda_core::tensor::{Shape, Metadata, TensorMetadata, strides};
use ruda_core::tensor::spatial::calculate_unfold_shape;
pub fn expand<R: Runtime>(tensor: RudaTensor<R>, target_shape: Shape) -> RudaTensor<R> {
if tensor.qparams.is_some() {
return super::expand_quantized::expand(tensor, target_shape);
}
let ndims_in = tensor.meta.shape().num_dims();
let ndims_out = target_shape.num_dims();
let mut new_strides = strides![0usize; ndims_out];
let dim_diff = ndims_out.saturating_sub(ndims_in);
let mut tensor_dim_iter = tensor.meta.shape().iter().rev();
for i in (0..ndims_out).rev() {
if i >= dim_diff {
if let Some(&tensor_dim) = tensor_dim_iter.next() {
if tensor_dim == target_shape[i] || tensor_dim == 1 {
new_strides[i] = if tensor_dim == target_shape[i] {
tensor.meta.strides()[i - dim_diff]
} else {
0
};
} else {
panic!(
"Dimension mismatch: cannot broadcast dimension {tensor_dim} of tensor to target shape"
);
}
} else {
new_strides[i] = 0;
}
} else {
new_strides[i] = 0;
}
}
RudaTensor {
client: tensor.client.clone(),
device: tensor.device.clone(),
meta: Box::new(Metadata::new(target_shape, new_strides)),
handle: tensor.handle.clone(),
dtype: tensor.dtype,
qparams: tensor.qparams.clone(),
}
}
pub fn unfold<R: Runtime>(
tensor: RudaTensor<R>,
dim: usize,
size: usize,
step: usize,
) -> RudaTensor<R> {
let shape = calculate_unfold_shape(tensor.shape(), dim, size, step);
let d_stride = tensor.meta.strides()[dim];
let mut strides = tensor.meta.strides.clone();
strides[dim] = step * d_stride;
strides.push(d_stride);
RudaTensor {
meta: Box::new(Metadata::new(shape, strides)),
client: tensor.client.clone(),
handle: tensor.handle.clone(),
device: tensor.device.clone(),
dtype: tensor.dtype,
qparams: tensor.qparams.clone(),
}
}