huggingface/ratchet

Improve slicing

Aperta

#91 aperta il 19 feb 2024

 (0 commenti) (0 reazioni) (0 assegnatari)Rust (44 fork)auto 404
enhancementhelp wanted

Metriche repository

Star
 (768 stelle)
Metriche merge PR
 (Metriche PR in attesa)

Descrizione

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.

Guida contributor