Skip to main content

gigastt_core/runtime/
tensor.rs

1/// Owned, cheaply cloneable tensor value used by the runtime abstraction layer.
2#[derive(Clone, Debug, PartialEq)]
3pub struct Tensor {
4    shape: Shape,
5    data: TensorData,
6}
7
8/// Owned tensor storage for the supported element types.
9#[derive(Clone, Debug, PartialEq)]
10pub enum TensorData {
11    F32(Vec<f32>),
12    I32(Vec<i32>),
13    I64(Vec<i64>),
14}
15
16/// Zero-copy borrow of tensor storage.
17#[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/// Normalized tensor shape independent of any runtime backend.
48#[derive(Clone, Debug, Default, PartialEq, Eq, Hash)]
49pub struct Shape {
50    dims: Vec<usize>,
51}
52
53/// Known tensor element types supported by the runtime abstraction.
54///
55/// Only types that have a corresponding [`TensorData`] variant are listed here;
56/// back-ends that produce other ONNX types must convert or reject them.
57#[derive(Clone, Copy, Debug, PartialEq, Eq)]
58pub enum ElementType {
59    F32,
60    I32,
61    I64,
62}
63
64impl Tensor {
65    /// Creates a tensor, validating that `data` length matches `shape`.
66    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    /// Convenience constructor that panics on shape/data mismatch.
79    ///
80    /// Use only when dimensions are statically known; prefer [`Self::new`] for
81    /// runtime-sized tensors.
82    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    /// Return a mutable view of the underlying f32 buffer, if this tensor is f32.
114    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    /// Return a mutable view of the underlying i32 buffer, if this tensor is i32.
122    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    /// Return a mutable view of the underlying i64 buffer, if this tensor is i64.
130    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    /// Resize the tensor to a new shape, reusing the existing storage.
138    ///
139    /// The buffer is resized to the new element count and zero-padded if it
140    /// grows. The caller must update the data before use.
141    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/// Borrowed view of a tensor.
167#[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}