gigastt_core/runtime/ort/
tensor.rs1use 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 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 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
55pub 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 #[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}