Skip to main content

onnx_extractor/
tensor.rs

1use prost::bytes::Bytes;
2use std::{mem, slice};
3
4use crate::{DataType, Error, external_data::ExternalDataInfo};
5
6#[derive(Debug, Clone)]
7pub(crate) enum TensorDataLocation {
8    None,
9    External(ExternalDataInfo),
10    Mmap(Bytes),
11    MmapStrings(Vec<Bytes>),
12    F32(Vec<f32>),
13    F64(Vec<f64>),
14    I64(Vec<i64>),
15    U64(Vec<u64>),
16    I32(Vec<i32>),
17}
18
19/// Zero-copy tensor data reference
20#[derive(Debug, Clone)]
21pub enum TensorDataRef<'a> {
22    Raw(Bytes),
23    Strings(&'a [Bytes]),
24    F32(&'a [f32]),
25    F64(&'a [f64]),
26    I32(&'a [i32]),
27    I64(&'a [i64]),
28    U64(&'a [u64]),
29}
30
31impl<'a> TensorDataRef<'a> {
32    /// Total byte length across all variants
33    ///
34    /// For `Strings`, returns the sum of all string element byte lengths.
35    /// If all `Strings` are empty, returns 0.
36    pub fn len(&self) -> usize {
37        match self {
38            TensorDataRef::Raw(b) => b.len(),
39            TensorDataRef::Strings(parts) => parts.iter().map(Bytes::len).sum(),
40            TensorDataRef::F32(v) => mem::size_of_val(*v),
41            TensorDataRef::F64(v) => mem::size_of_val(*v),
42            TensorDataRef::I32(v) => mem::size_of_val(*v),
43            TensorDataRef::I64(v) => mem::size_of_val(*v),
44            TensorDataRef::U64(v) => mem::size_of_val(*v),
45        }
46    }
47
48    /// Returns true if data contains no elements
49    ///
50    /// For `Raw` and numeric variants, equivalent to `len() == 0`.
51    /// For `Strings`, checks if the slice of string elements is empty.
52    pub fn is_empty(&self) -> bool {
53        match self {
54            TensorDataRef::Raw(b) => b.is_empty(),
55            TensorDataRef::Strings(s) => s.is_empty(),
56            TensorDataRef::F32(v) => v.is_empty(),
57            TensorDataRef::F64(v) => v.is_empty(),
58            TensorDataRef::I32(v) => v.is_empty(),
59            TensorDataRef::I64(v) => v.is_empty(),
60            TensorDataRef::U64(v) => v.is_empty(),
61        }
62    }
63
64    /// Get data as contiguous byte slice
65    ///
66    /// Raw and numeric variants borrow directly.
67    /// Returns `None` for the `Strings` variant as string arrays are not contiguous byte buffers.
68    pub fn as_slice(&self) -> Option<&[u8]> {
69        match self {
70            TensorDataRef::Raw(b) => Some(b),
71            TensorDataRef::F32(v) => Some(slice_as_u8(v)),
72            TensorDataRef::F64(v) => Some(slice_as_u8(v)),
73            TensorDataRef::I32(v) => Some(slice_as_u8(v)),
74            TensorDataRef::I64(v) => Some(slice_as_u8(v)),
75            TensorDataRef::U64(v) => Some(slice_as_u8(v)),
76            TensorDataRef::Strings(_) => None,
77        }
78    }
79
80    /// Access the string elements if the variant is `Strings`. Returns `None` otherwise.
81    pub fn strings(&self) -> Option<&'a [Bytes]> {
82        match self {
83            TensorDataRef::Strings(v) => Some(v),
84            _ => None,
85        }
86    }
87}
88
89/// Zero-copy owned tensor data
90#[derive(Debug, Clone)]
91pub enum TensorData {
92    Raw(Bytes),
93    Strings(Vec<Bytes>),
94    F32(Vec<f32>),
95    F64(Vec<f64>),
96    I32(Vec<i32>),
97    I64(Vec<i64>),
98    U64(Vec<u64>),
99}
100
101impl TensorData {
102    /// Total byte length across all variants
103    ///
104    /// For `Strings`, returns the sum of all string element byte lengths.
105    /// If all `Strings` are empty, returns 0.
106    pub fn len(&self) -> usize {
107        match self {
108            TensorData::Raw(b) => b.len(),
109            TensorData::Strings(parts) => parts.iter().map(Bytes::len).sum(),
110            TensorData::F32(v) => mem::size_of_val(v.as_slice()),
111            TensorData::F64(v) => mem::size_of_val(v.as_slice()),
112            TensorData::I32(v) => mem::size_of_val(v.as_slice()),
113            TensorData::I64(v) => mem::size_of_val(v.as_slice()),
114            TensorData::U64(v) => mem::size_of_val(v.as_slice()),
115        }
116    }
117
118    /// Returns true if data contains no elements
119    ///
120    /// For `Raw` and numeric variants, equivalent to `len() == 0`.
121    /// For `Strings`, checks if the vector of string elements is empty.
122    pub fn is_empty(&self) -> bool {
123        match self {
124            TensorData::Raw(b) => b.is_empty(),
125            TensorData::Strings(s) => s.is_empty(),
126            TensorData::F32(v) => v.is_empty(),
127            TensorData::F64(v) => v.is_empty(),
128            TensorData::I32(v) => v.is_empty(),
129            TensorData::I64(v) => v.is_empty(),
130            TensorData::U64(v) => v.is_empty(),
131        }
132    }
133
134    /// Get data as contiguous byte slice
135    ///
136    /// Raw and numeric variants borrow directly.
137    /// Returns `None` for the `Strings` variant as string arrays are not contiguous byte buffers.
138    pub fn as_slice(&self) -> Option<&[u8]> {
139        match self {
140            TensorData::Raw(b) => Some(b),
141            TensorData::F32(v) => Some(slice_as_u8(v)),
142            TensorData::F64(v) => Some(slice_as_u8(v)),
143            TensorData::I32(v) => Some(slice_as_u8(v)),
144            TensorData::I64(v) => Some(slice_as_u8(v)),
145            TensorData::U64(v) => Some(slice_as_u8(v)),
146            TensorData::Strings(_) => None,
147        }
148    }
149
150    /// Access the string elements if the variant is `Strings`. Returns `None` otherwise.
151    pub fn strings(&self) -> Option<&[Bytes]> {
152        match self {
153            TensorData::Strings(v) => Some(v),
154            _ => None,
155        }
156    }
157}
158
159/// An ONNX tensor with a name, shape, data type, and optional underlying data
160#[derive(Debug)]
161pub struct Tensor {
162    name: Option<String>,
163    shape: Vec<i64>,
164    data_type: DataType,
165    data: TensorDataLocation,
166}
167
168impl Tensor {
169    pub(crate) fn new(
170        name: Option<String>,
171        shape: Vec<i64>,
172        data_type: DataType,
173        data: TensorDataLocation,
174    ) -> Self {
175        Tensor {
176            name,
177            shape,
178            data_type,
179            data,
180        }
181    }
182
183    /// Tensor name
184    pub fn name(&self) -> Option<&str> {
185        self.name.as_deref()
186    }
187
188    /// Tensor shape dimensions
189    pub fn shape(&self) -> &[i64] {
190        &self.shape
191    }
192
193    /// Tensor data type
194    pub fn data_type(&self) -> DataType {
195        self.data_type
196    }
197
198    /// Returns true if this tensor contains data.
199    ///
200    /// This check does not trigger loading or memory-mapping of external files.
201    pub fn has_data(&self) -> bool {
202        !matches!(self.data, TensorDataLocation::None)
203    }
204
205    /// Borrow tensor data
206    ///
207    /// - All variants are returned without copying the underlying tensor data.
208    /// - External data is loaded from disk if not already in memory.
209    pub fn data(&self) -> Result<TensorDataRef<'_>, Error> {
210        match &self.data {
211            TensorDataLocation::External(external_info) => {
212                Ok(TensorDataRef::Raw(external_info.load_data()?))
213            }
214            TensorDataLocation::Mmap(bytes) => Ok(TensorDataRef::Raw(bytes.clone())),
215            TensorDataLocation::MmapStrings(strings) => Ok(TensorDataRef::Strings(strings)),
216            TensorDataLocation::F32(v) => Ok(TensorDataRef::F32(v)),
217            TensorDataLocation::F64(v) => Ok(TensorDataRef::F64(v)),
218            TensorDataLocation::I64(v) => Ok(TensorDataRef::I64(v)),
219            TensorDataLocation::U64(v) => Ok(TensorDataRef::U64(v)),
220            TensorDataLocation::I32(v) => Ok(TensorDataRef::I32(v)),
221            TensorDataLocation::None => Err(Error::MissingField("tensor data")),
222        }
223    }
224
225    /// Consume tensor and return owned data
226    ///
227    /// - All variants are returned without copying the underlying tensor data.
228    /// - External data is loaded from disk if not already in memory.
229    pub fn into_data(self) -> Result<TensorData, Error> {
230        match self.data {
231            TensorDataLocation::External(external_info) => {
232                Ok(TensorData::Raw(external_info.load_data()?))
233            }
234            TensorDataLocation::Mmap(bytes) => Ok(TensorData::Raw(bytes)),
235            TensorDataLocation::MmapStrings(strings) => Ok(TensorData::Strings(strings)),
236            TensorDataLocation::F32(v) => Ok(TensorData::F32(v)),
237            TensorDataLocation::F64(v) => Ok(TensorData::F64(v)),
238            TensorDataLocation::I64(v) => Ok(TensorData::I64(v)),
239            TensorDataLocation::U64(v) => Ok(TensorData::U64(v)),
240            TensorDataLocation::I32(v) => Ok(TensorData::I32(v)),
241            TensorDataLocation::None => Err(Error::MissingField("tensor data")),
242        }
243    }
244}
245
246fn slice_as_u8<T>(slice: &[T]) -> &[u8] {
247    unsafe { slice::from_raw_parts(slice.as_ptr().cast::<u8>(), mem::size_of_val(slice)) }
248}