burn-backend 0.22.0-pre.1

Core backend interfaces and data structures for executing tensor operations in Burn.
Documentation
use crate::{Backend, TensorMetadata, tensor::Device};
use alloc::vec::Vec;
use burn_std::{DType, Shape, Slice};

pub(crate) fn repeat_with_slice_assign<B, T, E, SA>(
    tensor: T,
    dim: usize,
    times: usize,
    device: Device<B>,
    empty: E,
    slice_assign: SA,
) -> T
where
    T: TensorMetadata,
    B: Backend,
    E: Fn(Shape, &Device<B>, DType) -> T,
    SA: Fn(T, &[Slice], T) -> T,
{
    let shape = tensor.shape();
    let dtype = tensor.dtype();

    let original_dim_length = shape[dim];
    let shape = shape.repeat(dim, times).unwrap();

    let mut tensor_output = empty(shape.clone(), &device, dtype);

    let indices_select_all = shape.iter().map(|d| 0..*d).collect::<Vec<_>>();

    let mut output_index = 0;
    for _ in 0..times {
        let mut indices = indices_select_all.clone();
        indices[dim] = output_index..output_index + original_dim_length;
        output_index += original_dim_length;

        // Convert ranges to Slice
        let slices: Vec<Slice> = indices
            .iter()
            .map(|r| Slice::new(r.start as isize, Some(r.end as isize), 1))
            .collect();
        tensor_output = slice_assign(tensor_output, &slices, tensor.clone());
    }

    tensor_output
}