ruprim_host/
repeat_dim.rs1use alloc::vec::Vec;
4use ruda_core::{bytes::Bytes, tensor::Shape};
5
6use ruda_core::tensor::host::{HostTensor, Layout};
7
8pub 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}