1#[derive(Clone, Debug, PartialEq)]
3pub struct Tensor {
4 shape: Shape,
5 data: TensorData,
6}
7
8#[derive(Clone, Debug, PartialEq)]
10pub enum TensorData {
11 F32(Vec<f32>),
12 I32(Vec<i32>),
13 I64(Vec<i64>),
14}
15
16#[derive(Clone, Copy, Debug, PartialEq)]
18pub enum TensorDataView<'a> {
19 F32(&'a [f32]),
20 I32(&'a [i32]),
21 I64(&'a [i64]),
22}
23
24impl<'a> TensorDataView<'a> {
25 pub fn as_f32(&self) -> Option<&'a [f32]> {
26 match self {
27 TensorDataView::F32(v) => Some(v),
28 _ => None,
29 }
30 }
31
32 pub fn as_i32(&self) -> Option<&'a [i32]> {
33 match self {
34 TensorDataView::I32(v) => Some(v),
35 _ => None,
36 }
37 }
38
39 pub fn as_i64(&self) -> Option<&'a [i64]> {
40 match self {
41 TensorDataView::I64(v) => Some(v),
42 _ => None,
43 }
44 }
45}
46
47#[derive(Clone, Debug, Default, PartialEq, Eq, Hash)]
49pub struct Shape {
50 dims: Vec<usize>,
51}
52
53#[derive(Clone, Copy, Debug, PartialEq, Eq)]
58pub enum ElementType {
59 F32,
60 I32,
61 I64,
62}
63
64impl Tensor {
65 pub fn new(shape: Shape, data: TensorData) -> Result<Self, crate::runtime::RuntimeError> {
67 let expected = shape.elements();
68 let actual = data.len();
69 if expected != actual {
70 return Err(crate::runtime::RuntimeError::DataLengthMismatch {
71 expected,
72 got: actual,
73 });
74 }
75 Ok(Self { shape, data })
76 }
77
78 pub fn new_checked(shape: Shape, data: TensorData) -> Self {
83 Self::new(shape, data).expect("tensor data length mismatch")
84 }
85
86 pub fn shape(&self) -> &Shape {
87 &self.shape
88 }
89
90 pub fn element_type(&self) -> ElementType {
91 match &self.data {
92 TensorData::F32(_) => ElementType::F32,
93 TensorData::I32(_) => ElementType::I32,
94 TensorData::I64(_) => ElementType::I64,
95 }
96 }
97
98 pub fn view(&self) -> TensorView<'_> {
99 TensorView {
100 shape: &self.shape,
101 data: match &self.data {
102 TensorData::F32(v) => TensorDataView::F32(v.as_slice()),
103 TensorData::I32(v) => TensorDataView::I32(v.as_slice()),
104 TensorData::I64(v) => TensorDataView::I64(v.as_slice()),
105 },
106 }
107 }
108
109 pub fn into_data(self) -> TensorData {
110 self.data
111 }
112
113 pub fn as_f32_mut(&mut self) -> Option<&mut [f32]> {
115 match &mut self.data {
116 TensorData::F32(v) => Some(v.as_mut_slice()),
117 _ => None,
118 }
119 }
120
121 pub fn as_i32_mut(&mut self) -> Option<&mut [i32]> {
123 match &mut self.data {
124 TensorData::I32(v) => Some(v.as_mut_slice()),
125 _ => None,
126 }
127 }
128
129 pub fn as_i64_mut(&mut self) -> Option<&mut [i64]> {
131 match &mut self.data {
132 TensorData::I64(v) => Some(v.as_mut_slice()),
133 _ => None,
134 }
135 }
136
137 pub fn resize_to(&mut self, shape: Shape) {
142 let new_len = shape.elements();
143 match &mut self.data {
144 TensorData::F32(v) => v.resize(new_len, 0.0),
145 TensorData::I32(v) => v.resize(new_len, 0),
146 TensorData::I64(v) => v.resize(new_len, 0),
147 }
148 self.shape = shape;
149 }
150}
151
152impl TensorData {
153 pub fn len(&self) -> usize {
154 match self {
155 TensorData::F32(v) => v.len(),
156 TensorData::I32(v) => v.len(),
157 TensorData::I64(v) => v.len(),
158 }
159 }
160
161 pub fn is_empty(&self) -> bool {
162 self.len() == 0
163 }
164}
165
166#[derive(Clone, Copy, Debug, PartialEq)]
168pub struct TensorView<'a> {
169 shape: &'a Shape,
170 data: TensorDataView<'a>,
171}
172
173impl<'a> TensorView<'a> {
174 pub fn shape(&self) -> &Shape {
175 self.shape
176 }
177
178 pub fn data(&self) -> &TensorDataView<'a> {
179 &self.data
180 }
181}
182
183impl Shape {
184 pub fn new(dims: Vec<usize>) -> Self {
185 Self { dims }
186 }
187
188 pub fn elements(&self) -> usize {
189 self.dims.iter().product()
190 }
191
192 pub fn dims(&self) -> &[usize] {
193 &self.dims
194 }
195}
196
197#[cfg(test)]
198mod tests {
199 use super::*;
200
201 #[test]
202 fn test_tensor_shape_and_data_match() {
203 let t = Tensor::new(Shape::new(vec![2, 3]), TensorData::F32(vec![0.0; 6])).unwrap();
204 assert_eq!(t.shape().dims(), &[2, 3]);
205 assert_eq!(t.element_type(), ElementType::F32);
206 }
207
208 #[test]
209 fn test_tensor_rejects_mismatched_data() {
210 let err = Tensor::new(Shape::new(vec![2, 3]), TensorData::F32(vec![0.0; 5])).unwrap_err();
211 assert!(matches!(
212 err,
213 crate::runtime::RuntimeError::DataLengthMismatch {
214 expected: 6,
215 got: 5
216 }
217 ));
218 }
219
220 #[test]
221 fn test_shape_elements() {
222 assert_eq!(Shape::new(vec![2, 3, 4]).elements(), 24);
223 assert_eq!(Shape::new(vec![]).elements(), 1);
224 }
225
226 #[test]
227 fn test_tensor_view_f32() {
228 let t = Tensor::new(
229 Shape::new(vec![2, 2]),
230 TensorData::F32(vec![1.0, 2.0, 3.0, 4.0]),
231 )
232 .unwrap();
233 let v = t.view();
234 assert_eq!(v.shape().dims(), &[2, 2]);
235 assert_eq!(v.data().as_f32(), Some(&[1.0, 2.0, 3.0, 4.0][..]));
236 }
237
238 #[test]
239 fn test_tensor_view_non_f32_returns_none() {
240 let t = Tensor::new(Shape::new(vec![2]), TensorData::I32(vec![1, 2])).unwrap();
241 let v = t.view();
242 assert_eq!(v.data().as_f32(), None);
243 }
244
245 #[test]
246 fn test_shape_elements_zero_dimension() {
247 assert_eq!(Shape::new(vec![0]).elements(), 0);
248 assert_eq!(Shape::new(vec![2, 0]).elements(), 0);
249 }
250}