candela/tensor/mem_formats/
slice.rs1use 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#[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#[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}