huggingface/ratchet

Improve slicing

Aberta

#91 aberto em 19 de fev. de 2024

 (0 comentário) (0 reação) (0 responsável)Rust (44 forks)auto 404
enhancementhelp wanted

Métricas do repositório

Stars
 (768 estrelas)
Métricas de merge de PR
 (Métricas PR pendentes)

Description

Currently our slice definition looks as follows:

    pub fn slice<D: std::ops::RangeBounds<usize>>(&self, ranges: &[D]) -> anyhow::Result<Tensor> {
        ///...impl...
    }

This is very user hostile, as the user must provide a homogeneous collection of ranges.

let y = x.slice(&[0..5, 0..6, 0..7]);

This is very annoying, as you may only care about one of the dimensions, it would be better to have something like

let y = x.slice(&[.., 0..6, ..]);

Unfortunately, this doesn't work, because the collection is now heterogeneous (i.e is made up of 2xRangeFull and 1xRange).

Therefore, we need to use a macro, much like ndarray.

let y = x.slice(s![.., 0..6, ..]);

This API gives the illusion of heterogeneous collections, which is what we want.

Guia do colaborador