use std::error::Error;
use rten::{Dimension, InputOrOutput, NodeId, Output};
#[derive(Clone)]
pub struct NodeInfo {
name: String,
shape: Vec<Dimension>,
}
impl NodeInfo {
pub fn name(&self) -> &str {
&self.name
}
pub fn shape(&self) -> &[Dimension] {
&self.shape
}
pub fn from_name_shape(name: &str, shape: &[Dimension]) -> NodeInfo {
NodeInfo {
name: name.to_string(),
shape: shape.to_vec(),
}
}
}
pub trait Model {
fn find_node(&self, name: &str) -> Option<NodeId>;
fn node_info(&self, id: NodeId) -> Option<NodeInfo>;
fn input_ids(&self) -> &[NodeId];
fn run(
&self,
inputs: Vec<(NodeId, InputOrOutput)>,
outputs: &[NodeId],
) -> Result<Vec<Output>, Box<dyn Error>>;
fn partial_run(
&self,
inputs: Vec<(NodeId, InputOrOutput)>,
outputs: &[NodeId],
) -> Result<Vec<(NodeId, Output)>, Box<dyn Error>>;
}
impl Model for rten::Model {
fn find_node(&self, name: &str) -> Option<NodeId> {
self.find_node(name)
}
fn node_info(&self, id: NodeId) -> Option<NodeInfo> {
self.node_info(id).and_then(|info| {
let name = info.name()?;
let dims = info.shape()?;
Some(NodeInfo {
name: name.to_string(),
shape: dims,
})
})
}
fn input_ids(&self) -> &[NodeId] {
self.input_ids()
}
fn run(
&self,
inputs: Vec<(NodeId, InputOrOutput)>,
outputs: &[NodeId],
) -> Result<Vec<Output>, Box<dyn Error>> {
self.run(inputs, outputs, None).map_err(|e| e.into())
}
fn partial_run(
&self,
inputs: Vec<(NodeId, InputOrOutput)>,
outputs: &[NodeId],
) -> Result<Vec<(NodeId, Output)>, Box<dyn Error>> {
self.partial_run(inputs, outputs, None)
.map_err(|e| e.into())
}
}