Skip to main content

rten_tensor/
slice_range.rs

1//! Range types used when slicing tensors.
2
3use smallvec::SmallVec;
4
5use std::fmt::Debug;
6use std::ops::{Range, RangeFrom, RangeFull, RangeTo};
7
8/// Specifies a subset of a dimension to include when slicing a tensor or view.
9///
10/// Can be constructed from an index or range using `index_or_range.into()`.
11#[derive(Clone, Copy, Debug, PartialEq)]
12pub enum SliceItem {
13    /// Extract a specific index from a dimension.
14    ///
15    /// The number of dimensions in the sliced view will be one minus the number
16    /// of dimensions sliced with an index. If the index is negative, it counts
17    /// back from the end of the dimension.
18    Index(isize),
19
20    /// Include a subset of the range of the dimension.
21    Range(SliceRange),
22}
23
24impl SliceItem {
25    /// Return a SliceItem that extracts the full range of a dimension.
26    #[inline]
27    pub fn full_range() -> Self {
28        (..).into()
29    }
30
31    /// Return a SliceItem that extracts part of an axis.
32    #[inline]
33    pub fn range(start: isize, end: Option<isize>, step: isize) -> SliceItem {
34        SliceItem::Range(SliceRange::new(start, end, step))
35    }
36
37    /// Return stepped index range selected by this item from an axis with a
38    /// given size.
39    pub(crate) fn index_range(&self, dim_size: usize) -> IndexRange {
40        let range = match *self {
41            SliceItem::Range(range) => range,
42            SliceItem::Index(idx) => SliceRange::new(idx, Some(idx + 1), 1),
43        };
44        range.index_range(dim_size)
45    }
46}
47
48// This conversion exists to avoid ambiguity when slicing a tensor with a
49// numeric literal of unspecified type (eg. `tensor.slice((0, 0))`). In this
50// case it is ambiguous which `SliceItem::from` should be used, but the i32
51// case is used if it exists.
52impl From<i32> for SliceItem {
53    #[inline]
54    fn from(value: i32) -> Self {
55        SliceItem::Index(value as isize)
56    }
57}
58
59impl From<isize> for SliceItem {
60    #[inline]
61    fn from(value: isize) -> Self {
62        SliceItem::Index(value)
63    }
64}
65
66impl From<usize> for SliceItem {
67    #[inline]
68    fn from(value: usize) -> Self {
69        SliceItem::Index(value as isize)
70    }
71}
72
73impl<R> From<R> for SliceItem
74where
75    R: Into<SliceRange>,
76{
77    fn from(value: R) -> Self {
78        SliceItem::Range(value.into())
79    }
80}
81
82/// Used to convert sequences of indices and/or ranges into a uniform
83/// `[SliceItem]` array that can be used to slice a tensor.
84///
85/// This trait is implemented for:
86///
87///  - Individual indices and ranges (types satisfying `Into<SliceItem>`)
88///  - Arrays of indices or ranges
89///  - Tuples of indices and/or ranges
90///  - `[SliceItem]` slices
91///
92/// Ranges can be specified using regular Rust ranges (eg. `start..end`,
93/// `start..`, `..end`, `..`) or a [`SliceRange`], which extends regular Rust
94/// ranges with support for steps and specifying endpoints using negative
95/// values, which behaves similarly to using negative values in NumPy.
96pub trait IntoSliceItems {
97    type Array: AsRef<[SliceItem]>;
98
99    fn into_slice_items(self) -> Self::Array;
100}
101
102impl<'a> IntoSliceItems for &'a [SliceItem] {
103    type Array = &'a [SliceItem];
104
105    fn into_slice_items(self) -> &'a [SliceItem] {
106        self
107    }
108}
109
110impl<const N: usize, T: Into<SliceItem>> IntoSliceItems for [T; N] {
111    type Array = [SliceItem; N];
112
113    fn into_slice_items(self) -> [SliceItem; N] {
114        self.map(|x| x.into())
115    }
116}
117
118impl<T: Into<SliceItem>> IntoSliceItems for T {
119    type Array = [SliceItem; 1];
120
121    fn into_slice_items(self) -> [SliceItem; 1] {
122        [self.into()]
123    }
124}
125
126impl<T1: Into<SliceItem>> IntoSliceItems for (T1,) {
127    type Array = [SliceItem; 1];
128
129    fn into_slice_items(self) -> [SliceItem; 1] {
130        [self.0.into()]
131    }
132}
133
134impl<T1: Into<SliceItem>, T2: Into<SliceItem>> IntoSliceItems for (T1, T2) {
135    type Array = [SliceItem; 2];
136
137    fn into_slice_items(self) -> [SliceItem; 2] {
138        [self.0.into(), self.1.into()]
139    }
140}
141
142impl<T1: Into<SliceItem>, T2: Into<SliceItem>, T3: Into<SliceItem>> IntoSliceItems
143    for (T1, T2, T3)
144{
145    type Array = [SliceItem; 3];
146
147    fn into_slice_items(self) -> [SliceItem; 3] {
148        [self.0.into(), self.1.into(), self.2.into()]
149    }
150}
151
152impl<T1: Into<SliceItem>, T2: Into<SliceItem>, T3: Into<SliceItem>, T4: Into<SliceItem>>
153    IntoSliceItems for (T1, T2, T3, T4)
154{
155    type Array = [SliceItem; 4];
156
157    fn into_slice_items(self) -> [SliceItem; 4] {
158        [self.0.into(), self.1.into(), self.2.into(), self.3.into()]
159    }
160}
161
162impl<
163    T1: Into<SliceItem>,
164    T2: Into<SliceItem>,
165    T3: Into<SliceItem>,
166    T4: Into<SliceItem>,
167    T5: Into<SliceItem>,
168> IntoSliceItems for (T1, T2, T3, T4, T5)
169{
170    type Array = [SliceItem; 5];
171
172    fn into_slice_items(self) -> [SliceItem; 5] {
173        [
174            self.0.into(),
175            self.1.into(),
176            self.2.into(),
177            self.3.into(),
178            self.4.into(),
179        ]
180    }
181}
182
183/// Dynamically sized array of [`SliceItem`]s, which avoids allocating in the
184/// common case where the length is small.
185pub type DynSliceItems = SmallVec<[SliceItem; 5]>;
186
187/// Convert a slice of indices into [`SliceItem`]s.
188///
189/// To convert indices of a statically known length to [`SliceItem`]s, use
190/// [`IntoSliceItems`] instead. This function is for the case when the length
191/// is not statically known, but is assumed to likely be small.
192pub fn to_slice_items<T: Clone + Into<SliceItem>>(index: &[T]) -> DynSliceItems {
193    index.iter().map(|x| x.clone().into()).collect()
194}
195
196/// A range for slicing a [`Tensor`](crate::Tensor) or [`NdTensor`](crate::NdTensor).
197///
198/// This has two main differences from [`Range`].
199///
200/// - A non-zero step between indices can be specified. The step can be negative,
201///   which means that the dimension should be traversed in reverse order.
202/// - The `start` and `end` indexes can also be negative, in which case they
203///   count backwards from the end of the array.
204///
205/// This system for specifying slicing and indexing follows NumPy, which in
206/// turn strongly influenced slicing in ONNX.
207#[derive(Clone, Copy, Debug, PartialEq)]
208pub struct SliceRange {
209    /// First index in range.
210    pub start: isize,
211
212    /// Last index (exclusive) in range, or None if the range extends to the
213    /// end of a dimension.
214    pub end: Option<isize>,
215
216    /// The steps between adjacent elements selected by this range. This
217    /// is private so this module can enforce the invariant that it is non-zero.
218    step: isize,
219}
220
221impl SliceRange {
222    /// Create a new range from `start` to `end`. The `start` index is inclusive
223    /// and the `end` value is exclusive. If `end` is None, the range spans
224    /// to the end of the dimension.
225    ///
226    /// Panics if the `step` size is 0.
227    #[inline]
228    pub fn new(start: isize, end: Option<isize>, step: isize) -> SliceRange {
229        assert!(step != 0, "Slice step cannot be 0");
230        SliceRange { start, end, step }
231    }
232
233    /// Return the number of elements that would be retained if using this range
234    /// to slice a dimension of size `dim_size`.
235    pub fn steps(&self, dim_size: usize) -> usize {
236        let clamped = self.clamp(dim_size);
237
238        let start_idx = Self::offset_from_start(clamped.start, dim_size);
239        let end_idx = clamped
240            .end
241            .map(|index| Self::offset_from_start(index, dim_size))
242            .unwrap_or(if self.step > 0 { dim_size as isize } else { -1 });
243
244        if (clamped.step > 0 && end_idx <= start_idx) || (clamped.step < 0 && end_idx >= start_idx)
245        {
246            return 0;
247        }
248
249        let steps = if clamped.step > 0 {
250            1 + (end_idx - start_idx - 1) / clamped.step
251        } else {
252            1 + (start_idx - end_idx - 1) / -clamped.step
253        };
254
255        steps.max(0) as usize
256    }
257
258    /// Return a copy of this range with indexes adjusted so that they are valid
259    /// for a tensor dimension of size `dim_size`.
260    ///
261    /// Valid indexes depend on direction that the dimension is traversed
262    /// (forwards if `self.step` is positive or backwards if negative). They
263    /// start at the first element going in that direction and end after the
264    /// last element.
265    pub fn clamp(&self, dim_size: usize) -> SliceRange {
266        let len = dim_size as isize;
267
268        let min_idx;
269        let max_idx;
270
271        if self.step > 0 {
272            // When traversing forwards, the range of valid +ve indexes is `[0,
273            // len]` and for -ve indexes `[-len, -1]`.
274            min_idx = -len;
275            max_idx = len;
276        } else {
277            // When traversing backwards, the range of valid +ve indexes are
278            // `[0, len-1]` and for -ve indexes `[-len-1, -1]`.
279            min_idx = -len - 1;
280            max_idx = len - 1;
281        }
282
283        SliceRange::new(
284            self.start.clamp(min_idx, max_idx),
285            self.end.map(|e| e.clamp(min_idx, max_idx)),
286            self.step,
287        )
288    }
289
290    pub fn step(&self) -> isize {
291        self.step
292    }
293
294    /// Clamp this range so that it is valid for a dimension of size `dim_size`
295    /// and resolve it to a positive range.
296    ///
297    /// This method is useful for implementing Python/NumPy-style slicing where
298    /// range endpoints can be out of bounds.
299    pub fn resolve_clamped(&self, dim_size: usize) -> Range<usize> {
300        self.clamp(dim_size).resolve(dim_size).unwrap()
301    }
302
303    /// Resolve the range endpoints to a positive range in `[0, dim_size)`.
304    ///
305    /// Returns the range if resolved or None if out of bounds.
306    ///
307    /// If `self.step` is positive, the returned range counts forwards from
308    /// the first index of the dimension, otherwise it counts backwards from
309    /// the last index.
310    #[inline]
311    pub fn resolve(&self, dim_size: usize) -> Option<Range<usize>> {
312        let (start, end) = if self.step > 0 {
313            let start = Self::offset_from_start(self.start, dim_size);
314            let end = self
315                .end
316                .map(|end| Self::offset_from_start(end, dim_size))
317                .unwrap_or(dim_size as isize);
318            (start, end)
319        } else {
320            let start = Self::offset_from_end(self.start, dim_size);
321            let end = self
322                .end
323                .map(|end| Self::offset_from_end(end, dim_size))
324                .unwrap_or(dim_size as isize);
325            (start, end)
326        };
327
328        if start >= 0 && start <= dim_size as isize && end >= 0 && end <= dim_size as isize {
329            // If `end < start` this means the range is empty. Set `end ==
330            // start` to have a canonical representation for this case.
331            let end = end.max(start);
332
333            Some(start as usize..end as usize)
334        } else {
335            None
336        }
337    }
338
339    /// Return stepped index range selected by this range from an axis with a
340    /// given size.
341    pub(crate) fn index_range(&self, dim_size: usize) -> IndexRange {
342        // Resolve range endpoints to `[0, N]`, counting forwards from the
343        // start if step > 0 or backwards from the end otherwise.
344        let resolved = self.resolve_clamped(dim_size);
345
346        if self.step > 0 {
347            IndexRange::new(resolved.start, resolved.end as isize, self.step)
348        } else {
349            IndexRange::new(
350                dim_size - 1 - resolved.start,
351                dim_size as isize - 1 - resolved.end as isize,
352                self.step,
353            )
354        }
355    }
356
357    /// Resolve an index to an offset from the first index of the dimension.
358    #[inline]
359    fn offset_from_start(index: isize, dim_size: usize) -> isize {
360        if index >= 0 {
361            index
362        } else {
363            dim_size as isize + index
364        }
365    }
366
367    /// Resolve an index to an offset from the last index of the dimension.
368    #[inline]
369    fn offset_from_end(index: isize, dim_size: usize) -> isize {
370        if index >= 0 {
371            dim_size as isize - 1 - index
372        } else {
373            -index - 1
374        }
375    }
376}
377
378impl<T> From<Range<T>> for SliceRange
379where
380    T: TryInto<isize>,
381    <T as TryInto<isize>>::Error: Debug,
382{
383    fn from(r: Range<T>) -> SliceRange {
384        let start = r.start.try_into().unwrap();
385        let end = r.end.try_into().unwrap();
386        SliceRange::new(start, Some(end), 1)
387    }
388}
389
390impl<T> From<RangeTo<T>> for SliceRange
391where
392    T: TryInto<isize>,
393    <T as TryInto<isize>>::Error: Debug,
394{
395    fn from(r: RangeTo<T>) -> SliceRange {
396        let end = r.end.try_into().unwrap();
397        SliceRange::new(0, Some(end), 1)
398    }
399}
400
401impl<T> From<RangeFrom<T>> for SliceRange
402where
403    T: TryInto<isize>,
404    <T as TryInto<isize>>::Error: Debug,
405{
406    fn from(r: RangeFrom<T>) -> SliceRange {
407        let start = r.start.try_into().unwrap();
408        SliceRange::new(start, None, 1)
409    }
410}
411
412impl From<RangeFull> for SliceRange {
413    #[inline]
414    fn from(_: RangeFull) -> SliceRange {
415        SliceRange::new(0, None, 1)
416    }
417}
418
419/// A range of indices with a step, which may be positive or negative.
420#[derive(Copy, Clone, Debug, PartialEq)]
421pub(crate) struct IndexRange {
422    /// Start index in [0, (dim_size - 1).max(0)]
423    start: usize,
424
425    /// End index in [-1, dim_size]
426    end: isize,
427    step: isize,
428}
429
430impl IndexRange {
431    /// Create a new range which steps from `start` (inclusive) to `end`
432    /// (exclusive) with a given step.
433    ///
434    /// The `step` value must not be zero.
435    ///
436    /// The `end` argument is signed to allow for a range which yields index 0
437    /// when `step` is negative. eg. `SteppedIndexRange::new(4, -1, -1)` will
438    /// yield indices `[4, 3, 2, 1, 0]`.
439    fn new(start: usize, end: isize, step: isize) -> Self {
440        assert!(step != 0);
441        assert!(start <= isize::MAX as usize);
442
443        IndexRange {
444            start,
445            end: end.max(-1),
446            step,
447        }
448    }
449
450    /// Return the start index.
451    #[allow(unused)]
452    pub fn start(&self) -> usize {
453        self.start
454    }
455
456    /// Return the index that is one past the end. This is signed since this
457    /// index can be -1 when `self.step() < 0`.
458    #[allow(unused)]
459    pub fn end(&self) -> isize {
460        self.end
461    }
462
463    /// Return the increment between indices.
464    #[allow(unused)]
465    pub fn step(&self) -> isize {
466        self.step
467    }
468
469    /// Return the number of steps along this dimension.
470    pub fn steps(&self) -> usize {
471        let len = if self.step > 0 {
472            (self.end - self.start as isize).max(0).unsigned_abs()
473        } else {
474            (self.end - self.start as isize).min(0).unsigned_abs()
475        };
476        len.div_ceil(self.step.unsigned_abs())
477    }
478}
479
480impl IntoIterator for IndexRange {
481    type Item = usize;
482    type IntoIter = IndexRangeIter;
483
484    #[inline]
485    fn into_iter(self) -> IndexRangeIter {
486        IndexRangeIter {
487            step: self.step,
488            index: self.start as isize,
489            remaining: self.steps(),
490        }
491    }
492}
493
494/// An iterator over the indices in an [`IndexRange`].
495#[derive(Clone, Debug, PartialEq)]
496pub(crate) struct IndexRangeIter {
497    /// Next index. This is in the range [-1, N] where `N` is the size of
498    /// the dimension. The values yielded by `next` are always in [0, N).
499    index: isize,
500
501    /// Remaining indices to yield.
502    remaining: usize,
503
504    step: isize,
505}
506
507impl Iterator for IndexRangeIter {
508    type Item = usize;
509
510    #[inline]
511    fn next(&mut self) -> Option<usize> {
512        if self.remaining == 0 {
513            return None;
514        }
515        let idx = self.index;
516        self.index += self.step;
517        self.remaining -= 1;
518        Some(idx as usize)
519    }
520
521    #[inline]
522    fn size_hint(&self) -> (usize, Option<usize>) {
523        (self.remaining, Some(self.remaining))
524    }
525}
526
527impl ExactSizeIterator for IndexRangeIter {}
528impl std::iter::FusedIterator for IndexRangeIter {}
529
530#[cfg(test)]
531mod tests {
532    use rten_testing::TestCases;
533
534    use super::{IntoSliceItems, SliceItem, SliceRange};
535
536    #[test]
537    fn test_into_slice_items() {
538        let x = (42).into_slice_items();
539        assert_eq!(x, [SliceItem::Index(42)]);
540
541        let x = (2..5).into_slice_items();
542        assert_eq!(x, [SliceItem::Range((2..5).into())]);
543
544        let x = (..5).into_slice_items();
545        assert_eq!(x, [SliceItem::Range((0..5).into())]);
546
547        let x = (3..).into_slice_items();
548        assert_eq!(x, [SliceItem::Range((3..).into())]);
549
550        let x = [1].into_slice_items();
551        assert_eq!(x, [SliceItem::Index(1)]);
552        let x = [1, 2].into_slice_items();
553        assert_eq!(x, [SliceItem::Index(1), SliceItem::Index(2)]);
554
555        let x = (0, 1..2, ..).into_slice_items();
556        assert_eq!(
557            x,
558            [
559                SliceItem::Index(0),
560                SliceItem::Range((1..2).into()),
561                SliceItem::full_range()
562            ]
563        );
564    }
565
566    #[test]
567    fn test_index_range() {
568        #[derive(Debug)]
569        struct Case {
570            range: SliceItem,
571            dim_size: usize,
572            indices: Vec<usize>,
573        }
574
575        let cases = [
576            // +ve step, +ve endpoints
577            Case {
578                range: SliceItem::range(0, Some(4), 1),
579                dim_size: 6,
580                indices: (0..4).collect(),
581            },
582            Case {
583                range: SliceItem::range(2, Some(4), 1),
584                dim_size: 6,
585                indices: vec![2, 3],
586            },
587            Case {
588                range: SliceItem::range(2, Some(128), 1),
589                dim_size: 5,
590                indices: vec![2, 3, 4],
591            },
592            // +ve step > 1, +ve endpoints
593            Case {
594                range: SliceItem::range(0, Some(5), 2),
595                dim_size: 5,
596                indices: vec![0, 2, 4],
597            },
598            // +ve step, no end
599            Case {
600                range: SliceItem::range(0, None, 1),
601                dim_size: 6,
602                indices: (0..6).collect(),
603            },
604            // +ve step, -ve endpoints
605            Case {
606                range: SliceItem::range(-1, Some(-6), 2),
607                dim_size: 5,
608                indices: vec![],
609            },
610            // -ve step, -ve endpoints
611            Case {
612                range: SliceItem::range(-1, Some(-128), -1),
613                dim_size: 5,
614                indices: vec![4, 3, 2, 1, 0],
615            },
616            // -ve step, no end
617            Case {
618                range: SliceItem::range(-1, None, -1),
619                dim_size: 5,
620                indices: vec![4, 3, 2, 1, 0],
621            },
622            // -ve step < -1, -ve endpoints
623            Case {
624                range: SliceItem::range(-1, Some(-6), -2),
625                dim_size: 5,
626                indices: vec![4, 2, 0],
627            },
628            // -ve step, +ve endpoints
629            Case {
630                range: SliceItem::range(1, Some(5), -2),
631                dim_size: 5,
632                indices: vec![],
633            },
634            // Empty range, +ve step
635            Case {
636                range: SliceItem::range(0, Some(0), 1),
637                dim_size: 4,
638                indices: vec![],
639            },
640            // Empty range, -ve step
641            Case {
642                range: SliceItem::range(0, Some(0), -1),
643                dim_size: 4,
644                indices: vec![],
645            },
646            // Single index
647            Case {
648                range: SliceItem::Index(2),
649                dim_size: 4,
650                indices: vec![2],
651            },
652            // Single index, out of range
653            Case {
654                range: SliceItem::Index(2),
655                dim_size: 0,
656                indices: vec![],
657            },
658        ];
659
660        cases.test_each(|case| {
661            let Case {
662                range,
663                dim_size,
664                indices,
665            } = case;
666
667            let mut index_iter = range.index_range(*dim_size).into_iter();
668            let size_hint = index_iter.size_hint();
669            let index_vec: Vec<_> = index_iter.by_ref().collect();
670
671            assert_eq!(size_hint, (index_vec.len(), Some(index_vec.len())));
672            assert_eq!(index_vec, *indices);
673            assert_eq!(index_iter.size_hint(), (0, Some(0)));
674        })
675    }
676
677    #[test]
678    fn test_index_range_steps() {
679        #[derive(Debug)]
680        struct Case {
681            range: SliceRange,
682            dim_size: usize,
683            steps: usize,
684        }
685
686        let cases = [
687            // Positive step, no end.
688            Case {
689                range: SliceRange::new(0, None, 1),
690                dim_size: 4,
691                steps: 4,
692            },
693            // Positive step size exceeds range length.
694            Case {
695                range: SliceRange::new(0, None, 5),
696                dim_size: 4,
697                steps: 1,
698            },
699            // Negative step, no end.
700            Case {
701                range: SliceRange::new(-1, None, -1),
702                dim_size: 3,
703                steps: 3,
704            },
705            // Negative step size exceeds range length.
706            Case {
707                range: SliceRange::new(1, Some(0), -2),
708                dim_size: 2,
709                steps: 1,
710            },
711        ];
712
713        cases.test_each(|case| {
714            assert_eq!(case.range.index_range(case.dim_size).steps(), case.steps);
715        })
716    }
717
718    #[test]
719    #[should_panic(expected = "Slice step cannot be 0")]
720    fn test_slice_range_zero_step() {
721        SliceRange::new(0, None, 0);
722    }
723
724    #[test]
725    fn test_slice_range_resolve() {
726        // +ve endpoints, +ve step
727        assert_eq!(SliceRange::new(0, Some(5), 1).resolve_clamped(10), 0..5);
728        assert_eq!(SliceRange::new(0, None, 1).resolve_clamped(10), 0..10);
729        assert_eq!(SliceRange::new(15, Some(20), 1).resolve_clamped(10), 10..10);
730        assert_eq!(SliceRange::new(15, Some(20), 1).resolve(10), None);
731        assert_eq!(SliceRange::new(4, None, 1).resolve(3), None);
732        assert_eq!(SliceRange::new(0, Some(10), 1).resolve(3), None);
733
734        // -ve endpoints, +ve step
735        assert_eq!(SliceRange::new(-5, Some(-1), 1).resolve_clamped(10), 5..9);
736        assert_eq!(SliceRange::new(-20, Some(-1), 1).resolve_clamped(10), 0..9);
737        assert_eq!(SliceRange::new(-20, Some(-1), 1).resolve(10), None);
738        assert_eq!(SliceRange::new(-5, None, 1).resolve_clamped(10), 5..10);
739
740        // +ve endpoints, -ve step.
741        //
742        // Note the returned ranges count backwards from the end of the
743        // dimension.
744        assert_eq!(SliceRange::new(5, Some(0), -1).resolve_clamped(10), 4..9);
745        assert_eq!(SliceRange::new(5, None, -1).resolve_clamped(10), 4..10);
746        assert_eq!(SliceRange::new(9, None, -1).resolve_clamped(10), 0..10);
747
748        // -ve endpoints, -ve step.
749        assert_eq!(SliceRange::new(-1, Some(-4), -1).resolve_clamped(3), 0..3);
750        assert_eq!(SliceRange::new(-1, None, -1).resolve_clamped(2), 0..2);
751    }
752}