Skip to main content

ruprim_host/
repeat_dim.rs

1//! Repeat a tensor along a dimension.
2
3use alloc::vec::Vec;
4use ruda_core::{bytes::Bytes, tensor::Shape};
5
6use ruda_core::tensor::host::{HostTensor, Layout};
7
8/// Repeat `tensor` along `dim` by `times`.
9pub fn repeat_dim(tensor: HostTensor, dim: usize, times: usize) -> HostTensor {
10    if times == 1 {
11        return tensor;
12    }
13
14    let ndims = tensor.layout().num_dims();
15    assert!(
16        dim < ndims,
17        "repeat_dim: dim {} out of bounds for tensor with {} dimensions",
18        dim,
19        ndims
20    );
21
22    let tensor = tensor.to_contiguous();
23    let shape = tensor.layout().shape().clone();
24    let dtype = tensor.dtype();
25    let elem_size = ruda_core::tensor::host::storage::dtype_size(dtype);
26
27    let mut new_dims: Vec<usize> = shape.iter().cloned().collect();
28    new_dims[dim] *= times;
29    let new_shape = Shape::from(new_dims);
30
31    let src: &[u8] = tensor.bytes();
32    let n = new_shape.num_elements() * elem_size;
33    let mut dst: Vec<u8> = Vec::with_capacity(n);
34
35    let inner: usize = shape.iter().skip(dim + 1).product();
36    let dim_size = shape[dim];
37    let chunk_bytes = dim_size * inner * elem_size;
38    let outer: usize = shape.iter().take(dim).product();
39
40    for o in 0..outer {
41        let start = o * chunk_bytes;
42        let end = start + chunk_bytes;
43        for _t in 0..times {
44            dst.extend_from_slice(&src[start..end]);
45        }
46    }
47
48    debug_assert_eq!(dst.len(), n);
49    HostTensor::new(
50        Bytes::from_bytes_vec(dst),
51        Layout::contiguous(new_shape),
52        dtype,
53    )
54}