Skip to main content

candela/tensor/mem_formats/
slice.rs

1use std::ops::{Range, RangeFrom, RangeFull, RangeTo};
2
3use crate::tensor::mem_formats::layout::Layout;
4
5use crate::tensor::errors::OpError;
6use crate::tensor::internals::calculate_adjacent_dim_stride;
7
8#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)]
9enum SliceBounds {
10    Beginning,
11    Index(usize),
12    ReverseIndex(usize),
13    End,
14}
15
16/// A per-axis range for [`slice`](crate::Tensor::slice), one entry per axis.
17///
18/// You rarely name `SliceRange` directly - the [`s!`](crate::s) macro builds the
19/// list from ordinary range syntax. It converts from `a..b`, `a..`, `..b`, `..`,
20/// and bare integers (a single index); negative bounds count from the end.
21///
22/// # Examples
23///
24/// ```
25/// use candela::{s, Dimension, Tensor};
26///
27/// let t = Tensor::from_slice(&[0.0, 1.0, 2.0, 3.0, 4.0, 5.0], &[2, 3]);
28/// // row 0, columns 1..3
29/// let sub = t.slice(s![0..1, 1..3])?.materialize();
30/// assert_eq!(sub.shape(), &[1, 2]);
31/// let vals: Vec<f64> = sub.iter().copied().collect();
32/// assert_eq!(vals, [1.0, 2.0]);
33/// # Ok::<(), candela::OpError>(())
34/// ```
35#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)]
36pub struct SliceRange {
37    start: SliceBounds,
38    end: SliceBounds,
39}
40
41impl From<i32> for SliceRange {
42    #[inline]
43    fn from(value: i32) -> Self {
44        if value >= 0 {
45            Self {
46                start: SliceBounds::Index(value as usize),
47                end: SliceBounds::Index((value + 1) as usize),
48            }
49        } else {
50            Self {
51                start: SliceBounds::ReverseIndex((-value) as usize),
52                end: SliceBounds::ReverseIndex((-(value + 1)) as usize),
53            }
54        }
55    }
56}
57
58impl From<RangeFrom<i32>> for SliceRange {
59    #[inline]
60    fn from(value: RangeFrom<i32>) -> Self {
61        if value.start >= 0 {
62            Self {
63                start: SliceBounds::Index(value.start as usize),
64                end: SliceBounds::End,
65            }
66        } else {
67            Self {
68                start: SliceBounds::ReverseIndex((-value.start) as usize),
69                end: SliceBounds::End,
70            }
71        }
72    }
73}
74
75impl From<RangeTo<i32>> for SliceRange {
76    #[inline]
77    fn from(value: RangeTo<i32>) -> Self {
78        if value.end >= 0 {
79            Self {
80                start: SliceBounds::Beginning,
81                end: SliceBounds::Index(value.end as usize),
82            }
83        } else {
84            Self {
85                start: SliceBounds::Beginning,
86                end: SliceBounds::ReverseIndex((-value.end) as usize),
87            }
88        }
89    }
90}
91
92impl From<RangeFull> for SliceRange {
93    #[inline]
94    fn from(_: RangeFull) -> Self {
95        Self {
96            start: SliceBounds::Beginning,
97            end: SliceBounds::End,
98        }
99    }
100}
101
102impl From<Range<i32>> for SliceRange {
103    #[inline]
104    fn from(value: Range<i32>) -> Self {
105        let start = if value.start >= 0 {
106            SliceBounds::Index(value.start as usize)
107        } else {
108            SliceBounds::ReverseIndex((-value.start) as usize)
109        };
110
111        let end = if value.end >= 0 {
112            SliceBounds::Index(value.end as usize)
113        } else {
114            SliceBounds::ReverseIndex((-value.end) as usize)
115        };
116
117        Self { start, end }
118    }
119}
120
121/////////////////////////////////////////////////////
122
123#[derive(Debug)]
124pub struct SliceInfo {
125    pub(crate) offset: usize,
126    pub(crate) shape: Box<[usize]>,
127    pub(crate) adj_stride: Box<[i32]>,
128}
129
130impl SliceInfo {
131    pub(crate) fn from_range(layout: &Layout, range: &[SliceRange]) -> Result<Self, OpError> {
132        if range.len() > layout.shape().len() {
133            return Err(OpError::AxesOutOfBounds);
134        }
135
136        let mut offset: i64 = layout.offset() as i64;
137        let mut new_shape: Vec<usize> = layout.shape().into();
138
139        for (dim, r) in range.iter().enumerate() {
140            let start = match r.start {
141                SliceBounds::Beginning => 0,
142                SliceBounds::Index(i) => {
143                    offset += (i as i64) * layout.stride()[dim] as i64;
144
145                    i
146                }
147                SliceBounds::ReverseIndex(i) => {
148                    let true_index = layout.shape()[dim] - i;
149                    offset += true_index as i64 * layout.stride()[dim] as i64;
150
151                    true_index
152                }
153                _ => unreachable!("a new variation of SliceBounds was implemented"),
154            };
155
156            let end = match r.end {
157                SliceBounds::End => layout.shape()[dim],
158                SliceBounds::Index(i) => i,
159                SliceBounds::ReverseIndex(i) => layout.shape()[dim] - i,
160                _ => unreachable!("a new variation of SliceBounds was implemented"),
161            };
162
163            if end <= start {
164                return Err(OpError::SliceOutOfBounds);
165            }
166
167            new_shape[dim] = end - start;
168        }
169
170        let len: usize = new_shape.iter().product();
171        if len + (offset as usize) > layout.len() {
172            return Err(OpError::InvalidSliceShape(layout.len(), len));
173        }
174
175        let adj_stride = calculate_adjacent_dim_stride(layout.stride(), &new_shape);
176
177        Ok(Self {
178            offset: offset as usize,
179            shape: new_shape.into_boxed_slice(),
180            adj_stride,
181        })
182    }
183}