huggingface/ratchet

Improve slicing

Ouverte

#91 ouverte le 19 févr. 2024

 (0 commentaire) (0 réaction) (0 personne assignée)Rust (44 forks)auto 404
enhancementhelp wanted

Métriques du dépôt

Stars
 (768 étoiles)
Métriques de merge PR
 (Métriques PR en attente)

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.

Guide contributeur