Skip to main content

gigastt_core/runtime/ort/
tensor.rs

1use ort::session::SessionInputValue;
2use ort::value::{TensorElementType, Value};
3
4use crate::runtime::{
5    error::RuntimeError,
6    tensor::{Shape, Tensor, TensorData, TensorDataView},
7};
8
9impl Tensor {
10    /// Converts this owned tensor into an `ort` value.
11    pub fn into_ort_value(self) -> Result<Value, RuntimeError> {
12        let shape: Vec<i64> = self.shape().dims().iter().map(|&d| d as i64).collect();
13        match self.into_data() {
14            TensorData::F32(data) => ort::value::Tensor::from_array((shape, data))
15                .map(|t| t.into_dyn())
16                .map_err(|e| RuntimeError::InferenceFailed(e.to_string())),
17            TensorData::I32(data) => ort::value::Tensor::from_array((shape, data))
18                .map(|t| t.into_dyn())
19                .map_err(|e| RuntimeError::InferenceFailed(e.to_string())),
20            TensorData::I64(data) => ort::value::Tensor::from_array((shape, data))
21                .map(|t| t.into_dyn())
22                .map_err(|e| RuntimeError::InferenceFailed(e.to_string())),
23        }
24    }
25
26    /// Returns a borrowed `ort` input value backed by this tensor's data.
27    ///
28    /// The returned `SessionInputValue` borrows from `self`; the caller must
29    /// keep this tensor alive for the duration of the `run` call.
30    pub fn as_ort_input(&self) -> Result<SessionInputValue<'_>, RuntimeError> {
31        let shape: Vec<i64> = self.shape().dims().iter().map(|&d| d as i64).collect();
32        match self.view().data() {
33            TensorDataView::F32(data) => {
34                let tensor_ref: ort::value::TensorRef<'_, f32> =
35                    ort::value::TensorRef::from_array_view((shape, *data))
36                        .map_err(|e| RuntimeError::InferenceFailed(e.to_string()))?;
37                Ok(tensor_ref.into_dyn().into())
38            }
39            TensorDataView::I32(data) => {
40                let tensor_ref: ort::value::TensorRef<'_, i32> =
41                    ort::value::TensorRef::from_array_view((shape, *data))
42                        .map_err(|e| RuntimeError::InferenceFailed(e.to_string()))?;
43                Ok(tensor_ref.into_dyn().into())
44            }
45            TensorDataView::I64(data) => {
46                let tensor_ref: ort::value::TensorRef<'_, i64> =
47                    ort::value::TensorRef::from_array_view((shape, *data))
48                        .map_err(|e| RuntimeError::InferenceFailed(e.to_string()))?;
49                Ok(tensor_ref.into_dyn().into())
50            }
51        }
52    }
53}
54
55/// Converts an `ort` tensor value into our owned tensor type.
56pub fn value_to_tensor(value: Value) -> Result<Tensor, RuntimeError> {
57    match *value.data_type() {
58        TensorElementType::Float32 => {
59            let (shape, data) = value
60                .try_extract_tensor::<f32>()
61                .map_err(|e| RuntimeError::InferenceFailed(e.to_string()))?;
62            Tensor::new(
63                Shape::new(shape.iter().map(|&d| d as usize).collect()),
64                TensorData::F32(data.to_vec()),
65            )
66        }
67        TensorElementType::Int32 => {
68            let (shape, data) = value
69                .try_extract_tensor::<i32>()
70                .map_err(|e| RuntimeError::InferenceFailed(e.to_string()))?;
71            Tensor::new(
72                Shape::new(shape.iter().map(|&d| d as usize).collect()),
73                TensorData::I32(data.to_vec()),
74            )
75        }
76        TensorElementType::Int64 => {
77            let (shape, data) = value
78                .try_extract_tensor::<i64>()
79                .map_err(|e| RuntimeError::InferenceFailed(e.to_string()))?;
80            Tensor::new(
81                Shape::new(shape.iter().map(|&d| d as usize).collect()),
82                TensorData::I64(data.to_vec()),
83            )
84        }
85        other => Err(RuntimeError::InferenceFailed(format!(
86            "unsupported element type: {other:?}"
87        ))),
88    }
89}
90
91#[cfg(test)]
92mod tests {
93    use super::*;
94
95    // Skipped under Miri: constructs a real `ort::value::Tensor`, which calls
96    // into the onnxruntime C API — a foreign function Miri cannot interpret.
97    #[test]
98    #[cfg_attr(miri, ignore = "calls into onnxruntime FFI")]
99    fn test_tensor_ort_roundtrip_f32() {
100        let tensor = Tensor::new(
101            Shape::new(vec![2, 3]),
102            TensorData::F32(vec![1.0, 2.0, 3.0, 4.0, 5.0, 6.0]),
103        )
104        .unwrap();
105        let value = tensor.clone().into_ort_value().unwrap();
106        let recovered = value_to_tensor(value).unwrap();
107        assert_eq!(tensor, recovered);
108    }
109
110    #[test]
111    #[cfg_attr(miri, ignore = "calls into onnxruntime FFI")]
112    fn test_tensor_ort_roundtrip_i32() {
113        let tensor = Tensor::new(Shape::new(vec![3]), TensorData::I32(vec![1, 2, 3])).unwrap();
114        let value = tensor.clone().into_ort_value().unwrap();
115        let recovered = value_to_tensor(value).unwrap();
116        assert_eq!(tensor, recovered);
117    }
118
119    #[test]
120    #[cfg_attr(miri, ignore = "calls into onnxruntime FFI")]
121    fn test_tensor_ort_roundtrip_i64() {
122        let tensor =
123            Tensor::new(Shape::new(vec![2, 2]), TensorData::I64(vec![1, 2, 3, 4])).unwrap();
124        let value = tensor.clone().into_ort_value().unwrap();
125        let recovered = value_to_tensor(value).unwrap();
126        assert_eq!(tensor, recovered);
127    }
128}