use std::fmt;
#[derive(Clone, Debug, PartialEq)]
pub enum TensorData {
F32(Vec<f32>),
I64(Vec<i64>),
}
#[derive(Clone, Debug, PartialEq)]
pub struct InferenceTensor {
pub shape: Vec<usize>,
pub data: TensorData,
}
impl InferenceTensor {
pub fn f32(shape: impl Into<Vec<usize>>, data: Vec<f32>) -> Self {
Self {
shape: shape.into(),
data: TensorData::F32(data),
}
}
pub fn i64(shape: impl Into<Vec<usize>>, data: Vec<i64>) -> Self {
Self {
shape: shape.into(),
data: TensorData::I64(data),
}
}
pub fn i64_scalar(value: i64) -> Self {
Self::i64(Vec::new(), vec![value])
}
pub fn as_f32_slice(&self) -> Result<&[f32], InferenceError> {
match &self.data {
TensorData::F32(v) => Ok(v.as_slice()),
TensorData::I64(_) => Err(InferenceError::TypeMismatch {
expected: "f32",
actual: "i64",
}),
}
}
pub fn into_f32(self) -> Result<Vec<f32>, InferenceError> {
match self.data {
TensorData::F32(v) => Ok(v),
TensorData::I64(_) => Err(InferenceError::TypeMismatch {
expected: "f32",
actual: "i64",
}),
}
}
}
#[derive(Clone, Debug)]
pub struct NamedTensor<'a> {
pub name: &'a str,
pub tensor: &'a InferenceTensor,
}
impl<'a> NamedTensor<'a> {
pub fn new(name: &'a str, tensor: &'a InferenceTensor) -> Self {
Self { name, tensor }
}
}
#[derive(Debug)]
pub enum InferenceError {
Load(String),
Run(String),
TypeMismatch {
expected: &'static str,
actual: &'static str,
},
MissingOutput { index: usize, available: usize },
}
impl fmt::Display for InferenceError {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
match self {
Self::Load(msg) => write!(f, "inference session load failed: {msg}"),
Self::Run(msg) => write!(f, "inference run failed: {msg}"),
Self::TypeMismatch { expected, actual } => {
write!(f, "tensor type mismatch: expected {expected}, got {actual}")
}
Self::MissingOutput { index, available } => {
write!(
f,
"missing output index {index} (model produced {available} outputs)"
)
}
}
}
}
impl std::error::Error for InferenceError {}
pub trait InferenceRuntime: Send {
fn input_names(&self) -> &[String];
fn output_names(&self) -> &[String] {
&[]
}
fn primary_input_name(&self) -> Option<&str> {
self.input_names().first().map(String::as_str)
}
fn run(&mut self, inputs: &[NamedTensor<'_>]) -> Result<Vec<InferenceTensor>, InferenceError>;
fn run_ordered(
&mut self,
inputs: &[&InferenceTensor],
) -> Result<Vec<InferenceTensor>, InferenceError>;
}
#[cfg(test)]
#[allow(clippy::unwrap_used)]
mod tests {
use super::*;
pub struct MockRuntime {
names: Vec<String>,
pub outputs: Vec<InferenceTensor>,
pub last_named: Vec<(String, InferenceTensor)>,
pub last_ordered: Vec<InferenceTensor>,
}
impl MockRuntime {
pub fn new(input_names: &[&str], outputs: Vec<InferenceTensor>) -> Self {
Self {
names: input_names.iter().map(|s| (*s).to_owned()).collect(),
outputs,
last_named: Vec::new(),
last_ordered: Vec::new(),
}
}
}
impl InferenceRuntime for MockRuntime {
fn input_names(&self) -> &[String] {
&self.names
}
fn run(
&mut self,
inputs: &[NamedTensor<'_>],
) -> Result<Vec<InferenceTensor>, InferenceError> {
self.last_named = inputs
.iter()
.map(|n| (n.name.to_owned(), n.tensor.clone()))
.collect();
Ok(self.outputs.clone())
}
fn run_ordered(
&mut self,
inputs: &[&InferenceTensor],
) -> Result<Vec<InferenceTensor>, InferenceError> {
self.last_ordered = inputs.iter().map(|t| (*t).clone()).collect();
Ok(self.outputs.clone())
}
}
#[test]
fn mock_runtime_round_trip_named() {
let out = InferenceTensor::f32(vec![1], vec![0.9]);
let mut rt = MockRuntime::new(&["input"], vec![out.clone()]);
let inp = InferenceTensor::f32(vec![1, 4], vec![0.0; 4]);
let got = rt
.run(&[NamedTensor::new("input", &inp)])
.expect("mock run");
assert_eq!(got, vec![out]);
assert_eq!(rt.last_named[0].0, "input");
}
#[test]
fn tensor_f32_accessors() {
let t = InferenceTensor::f32(vec![2, 2], vec![1.0, 2.0, 3.0, 4.0]);
assert_eq!(t.as_f32_slice().unwrap(), &[1.0, 2.0, 3.0, 4.0]);
let i = InferenceTensor::i64_scalar(16_000);
assert!(i.as_f32_slice().is_err());
}
#[test]
fn inference_error_display() {
let e = InferenceError::MissingOutput {
index: 1,
available: 1,
};
assert!(e.to_string().contains("missing output"));
}
}