use std::collections::HashMap;
use std::env;
use std::path::Path;
use std::sync::Arc;
#[cfg(feature = "mmap")]
use std::fs::File;
#[cfg(feature = "mmap")]
use memmap2::Mmap;
use crate::constant_storage::ConstantStorage;
use crate::env::str_as_bool;
use crate::graph::{Dimension, Graph, Node, NodeId, RunError, RunErrorImpl, RunOptions};
use crate::infer_shapes::InferShapeOptions;
use crate::op_registry::OpRegistry;
use crate::optimize::OptimizeOptions;
use crate::timing::{TimingFilter, TimingSort};
use crate::value::{Value, ValueOrView, ValueType};
use crate::weight_cache::WeightCache;
#[cfg(feature = "onnx_format")]
mod external_data;
mod file_type;
mod load_error;
mod metadata;
#[cfg(feature = "onnx_format")]
pub(crate) mod onnx_loader;
#[cfg(feature = "rten_format")]
mod rten_loader;
pub use load_error::{LoadError, LoadErrorKind};
pub use metadata::ModelMetadata;
use file_type::FileType;
use load_error::LoadErrorImpl;
#[cfg(test)]
pub mod rten_builder;
#[cfg(all(test, feature = "onnx_format"))]
pub mod onnx_builder;
pub struct Model {
graph: Graph,
metadata: ModelMetadata,
weight_cache: WeightCache,
}
impl Model {
pub fn load_file<P: AsRef<Path>>(path: P) -> Result<Model, LoadError> {
ModelOptions::with_all_ops().load_file(path)
}
pub fn load(data: Vec<u8>) -> Result<Model, LoadError> {
ModelOptions::with_all_ops().load(data)
}
pub fn load_static_slice(data: &'static [u8]) -> Result<Model, LoadError> {
ModelOptions::with_all_ops().load_static_slice(data)
}
#[cfg(feature = "mmap")]
#[cfg(not(target_arch = "wasm32"))]
pub unsafe fn load_mmap<P: AsRef<Path>>(path: P) -> Result<Model, LoadError> {
let opts = ModelOptions::with_all_ops();
unsafe { opts.load_mmap(path) }
}
pub fn find_node(&self, id: &str) -> Option<NodeId> {
self.graph.get_node_id(id)
}
pub fn node_id(&self, id: &str) -> Result<NodeId, RunError> {
self.find_node(id)
.ok_or_else(|| RunErrorImpl::InvalidNodeName(id.to_string()).into())
}
pub fn node_info(&self, id: NodeId) -> Option<NodeInfo<'_>> {
self.graph.get_node(id).map(|node| NodeInfo { node })
}
pub fn metadata(&self) -> &ModelMetadata {
&self.metadata
}
pub fn input_ids(&self) -> &[NodeId] {
self.graph.input_ids()
}
pub fn output_ids(&self) -> &[NodeId] {
self.graph.output_ids()
}
pub fn total_params(&self) -> usize {
self.graph.total_params()
}
pub fn input_shape(&self, index: usize) -> Option<Vec<Dimension>> {
let input_id = self.graph.input_ids().get(index)?;
let node_info = self.node_info(*input_id)?;
node_info.shape()
}
pub fn run(
&self,
inputs: Vec<(NodeId, ValueOrView)>,
outputs: &[NodeId],
opts: Option<RunOptions>,
) -> Result<Vec<Value>, RunError> {
let mut opts = opts.unwrap_or_default();
if let Some(timing_var) = env::var_os("RTEN_TIMING") {
let timing_var = timing_var.to_string_lossy();
parse_timing_config(&timing_var, &mut opts);
}
self.graph
.run(inputs, outputs, Some(&self.weight_cache), Some(opts))
}
pub fn run_n<const N: usize>(
&self,
inputs: Vec<(NodeId, ValueOrView)>,
outputs: [NodeId; N],
opts: Option<RunOptions>,
) -> Result<[Value; N], RunError> {
let result = self.run(inputs, &outputs, opts)?;
Ok(result.try_into().expect("wrong output count"))
}
pub fn run_one(&self, input: ValueOrView, opts: Option<RunOptions>) -> Result<Value, RunError> {
let &input_id = self
.input_ids()
.first()
.ok_or(RunErrorImpl::InvalidNodeId)?;
let &output_id = self
.output_ids()
.first()
.ok_or(RunErrorImpl::InvalidNodeId)?;
self.run_n(vec![(input_id, input)], [output_id], opts)
.map(|[result]| result)
}
pub fn partial_run(
&self,
inputs: Vec<(NodeId, ValueOrView)>,
outputs: &[NodeId],
opts: Option<RunOptions>,
) -> Result<Vec<(NodeId, Value)>, RunError> {
self.graph.partial_run(inputs, outputs, opts)
}
#[cfg(test)]
fn graph(&self) -> &Graph {
&self.graph
}
}
impl std::fmt::Debug for Model {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
let node_names = |ids: &[NodeId]| -> Vec<&str> {
ids.iter()
.filter_map(|id| self.node_info(*id))
.map(|info| info.name().unwrap_or(""))
.collect()
};
let input_names = node_names(self.input_ids());
let output_names = node_names(self.output_ids());
f.debug_struct("Model")
.field("inputs", &input_names)
.field("outputs", &output_names)
.finish()
}
}
pub struct NodeInfo<'a> {
node: &'a Node,
}
impl<'a> NodeInfo<'a> {
pub fn name(&self) -> Option<&'a str> {
self.node.name()
}
pub fn shape(&self) -> Option<Vec<Dimension>> {
self.node.shape().map(|n| n.into_owned())
}
pub fn dtype(&self) -> Option<ValueType> {
self.node.dtype()
}
}
impl<'a> std::fmt::Debug for NodeInfo<'a> {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("NodeInfo")
.field("name", &self.name())
.field("shape", &self.shape())
.field("dtype", &self.dtype())
.finish()
}
}
fn parse_timing_config(config: &str, opts: &mut RunOptions) {
opts.timing = true;
for token in config.split_ascii_whitespace() {
if let Some((key, val)) = token.split_once('=') {
let (key, val) = (key.trim(), val.trim());
match key {
"by-shape" => opts.timing_by_shape = str_as_bool(val),
"filter-op" => {
for op_name in val.split(',') {
opts.timing_filter
.push(TimingFilter::Operator(op_name.to_string()));
}
}
"sort" => match val {
"name" => opts.timing_sort = TimingSort::ByName,
"time" => opts.timing_sort = TimingSort::ByTime,
_ => eprintln!("Unrecognized sort order \"{}\"", val),
},
_ => {
eprintln!("Unrecognized timing option \"{}\"", key);
}
}
}
}
}
#[derive(Clone, Debug, PartialEq)]
pub enum ShapeInferenceMode {
Off,
On,
Strict,
}
#[derive(Clone)]
pub struct ModelOptions {
registry: Arc<OpRegistry>,
optimize: bool,
prepack_weights: bool,
external_data: HashMap<String, Arc<ConstantStorage>>,
infer_shapes: ShapeInferenceMode,
}
impl ModelOptions {
pub fn with_all_ops() -> ModelOptions {
Self::with_ops(OpRegistry::with_all_ops())
}
pub fn with_ops(ops: OpRegistry) -> ModelOptions {
ModelOptions {
registry: ops.into(),
optimize: true,
prepack_weights: false,
external_data: HashMap::new(),
infer_shapes: ShapeInferenceMode::On,
}
}
pub fn enable_optimization(&mut self, enable: bool) -> &mut Self {
self.optimize = enable;
self
}
#[deprecated]
pub fn enable_shape_inference(&mut self, enable: bool) -> &mut Self {
self.infer_shapes = if enable {
ShapeInferenceMode::On
} else {
ShapeInferenceMode::Off
};
self
}
pub fn shape_inference(&mut self, mode: ShapeInferenceMode) -> &mut Self {
self.infer_shapes = mode;
self
}
pub fn prepack_weights(&mut self, prepack: bool) -> &mut Self {
self.prepack_weights = prepack;
self
}
pub fn external_data(&mut self, path: &str, buf: Vec<u8>) -> &mut Self {
let storage = Arc::new(ConstantStorage::Buffer(buf));
self.external_data.insert(path.to_string(), storage);
self
}
pub fn load_file<P: AsRef<Path>>(&self, path: P) -> Result<Model, LoadError> {
match FileType::from_path(path.as_ref()).ok_or(LoadErrorImpl::UnknownFileType)? {
#[cfg(feature = "rten_format")]
FileType::Rten => {
use crate::constant_storage::ConstantStorage;
let data = std::fs::read(&path).map_err(LoadErrorImpl::ReadFailed)?;
let storage = Arc::new(ConstantStorage::Buffer(data));
rten_loader::load(storage, self)
}
#[cfg(not(feature = "rten_format"))]
FileType::Rten => Err(LoadErrorImpl::FormatNotEnabled.into()),
#[cfg(feature = "onnx_format")]
FileType::Onnx => {
let loader = external_data::FileLoader::new(path.as_ref())?;
onnx_loader::load(
onnx_loader::Source::Path(path.as_ref()),
Some(&loader),
self,
)
}
#[cfg(not(feature = "onnx_format"))]
FileType::Onnx => Err(LoadErrorImpl::FormatNotEnabled.into()),
}
}
#[cfg(feature = "onnx_format")]
fn mem_data_loader(&self) -> external_data::MemLoader {
let external_data = self.external_data.clone();
external_data::MemLoader::new(external_data)
}
pub fn load(&self, data: Vec<u8>) -> Result<Model, LoadError> {
match FileType::from_buffer(&data).ok_or(LoadErrorImpl::UnknownFileType)? {
#[cfg(feature = "rten_format")]
FileType::Rten => {
use crate::constant_storage::ConstantStorage;
let storage = Arc::new(ConstantStorage::Buffer(data));
rten_loader::load(storage, self)
}
#[cfg(not(feature = "rten_format"))]
FileType::Rten => Err(LoadErrorImpl::FormatNotEnabled.into()),
#[cfg(feature = "onnx_format")]
FileType::Onnx => {
let loader = self.mem_data_loader();
onnx_loader::load(onnx_loader::Source::Buffer(&data), Some(&loader), self)
}
#[cfg(not(feature = "onnx_format"))]
FileType::Onnx => Err(LoadErrorImpl::FormatNotEnabled.into()),
}
}
pub fn load_static_slice(&self, data: &'static [u8]) -> Result<Model, LoadError> {
match FileType::from_buffer(data).ok_or(LoadErrorImpl::UnknownFileType)? {
#[cfg(feature = "rten_format")]
FileType::Rten => {
use crate::constant_storage::ConstantStorage;
let storage = Arc::new(ConstantStorage::StaticSlice(data));
rten_loader::load(storage, self)
}
#[cfg(not(feature = "rten_format"))]
FileType::Rten => Err(LoadErrorImpl::FormatNotEnabled.into()),
#[cfg(feature = "onnx_format")]
FileType::Onnx => {
let loader = self.mem_data_loader();
onnx_loader::load(onnx_loader::Source::Buffer(data), Some(&loader), self)
}
#[cfg(not(feature = "onnx_format"))]
FileType::Onnx => Err(LoadErrorImpl::FormatNotEnabled.into()),
}
}
#[cfg(feature = "mmap")]
pub unsafe fn load_mmap<P: AsRef<Path>>(&self, path: P) -> Result<Model, LoadError> {
let file = File::open(&path).map_err(LoadErrorImpl::ReadFailed)?;
let mmap = unsafe { Mmap::map(&file) }.map_err(LoadErrorImpl::ReadFailed)?;
match FileType::from_path(path.as_ref()).ok_or(LoadErrorImpl::UnknownFileType)? {
#[cfg(feature = "rten_format")]
FileType::Rten => {
use crate::constant_storage::ConstantStorage;
let storage = Arc::new(ConstantStorage::Mmap(mmap));
rten_loader::load(storage, self)
}
#[cfg(not(feature = "rten_format"))]
FileType::Rten => Err(LoadErrorImpl::FormatNotEnabled.into()),
#[cfg(feature = "onnx_format")]
FileType::Onnx => {
let loader = unsafe { external_data::MmapLoader::new(path.as_ref()) }?;
onnx_loader::load(onnx_loader::Source::Buffer(&mmap), Some(&loader), self)
}
#[cfg(not(feature = "onnx_format"))]
FileType::Onnx => Err(LoadErrorImpl::FormatNotEnabled.into()),
}
}
fn optimize_mode(&self) -> OptimizeMode {
if self.optimize {
OptimizeMode::On(OptimizeOptions {
infer_shapes: match self.infer_shapes {
ShapeInferenceMode::Off => None,
ShapeInferenceMode::On => Some(InferShapeOptions {
strict: false,
..Default::default()
}),
ShapeInferenceMode::Strict => Some(InferShapeOptions {
strict: true,
..Default::default()
}),
},
})
} else {
OptimizeMode::Off
}
}
}
impl std::fmt::Debug for ModelOptions {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("ModelOptions")
.field("optimize", &self.optimize)
.field("prepack_weights", &self.prepack_weights)
.finish()
}
}
impl Default for ModelOptions {
fn default() -> Self {
ModelOptions::with_all_ops()
}
}
#[derive(Clone)]
enum OptimizeMode {
Off,
On(OptimizeOptions),
}
#[cfg(test)]
mod tests {
use rten_tensor::prelude::*;
use rten_tensor::{NdTensor, Tensor};
use crate::graph::{Dimension, NodeId, RunErrorKind};
use crate::model::rten_builder::{
GraphBuilder, IfArgs, MetadataArgs, ModelBuilder, ModelFormat, OpType,
};
use crate::model::{LoadErrorKind, Model, ModelOptions};
use crate::op_registry;
use crate::ops;
use crate::ops::{
BoxOrder, CoordTransformMode, DepthToSpaceMode, NearestMode, ResizeMode, Shape,
};
use crate::value::{DataType, Scalar, Value, ValueType};
fn generate_model_buffer(format: ModelFormat) -> Vec<u8> {
let mut builder = ModelBuilder::new(format);
let mut graph_builder = builder.graph_builder();
let const_val = Tensor::from_data(&[1, 2, 2], vec![0.5, -0.5, 0.1, -0.1]);
let const_node = graph_builder.add_constant(const_val.view());
let input_shape: Vec<Dimension> = const_val
.shape()
.iter()
.copied()
.map(Dimension::Fixed)
.collect();
let input_node =
graph_builder.add_value("input", Some(&input_shape), Some(DataType::Float));
let output_node = graph_builder.add_value("output", None, Some(DataType::Float));
graph_builder.add_input(input_node);
graph_builder.add_output(output_node);
let concat_out = graph_builder.add_value("concat_out", None, None);
graph_builder.add_operator(
"concat",
OpType::Concat(ops::Concat { axis: 0 }),
&[const_node, input_node].map(Some),
&[concat_out],
);
graph_builder.add_operator("relu", OpType::Relu, &[Some(concat_out)], &[output_node]);
let graph = graph_builder.finish();
builder.set_graph(graph);
builder.add_metadata(MetadataArgs {
onnx_hash: Some("abc".to_string()),
});
builder.finish()
}
fn generate_input() -> Tensor<f32> {
Tensor::from_data(&[1, 2, 2], vec![1., 2., -1., -2.])
}
fn check_output(mut result: Vec<Value>) -> Tensor<f32> {
assert_eq!(result.len(), 1);
let tensor: Tensor<f32> = result.remove(0).into_tensor::<f32>().unwrap();
assert_eq!(tensor.shape(), &[2, 2, 2]);
assert_eq!(tensor.to_vec(), &[0.5, 0., 0.1, 0., 1., 2., 0., 0.]);
tensor
}
#[test]
fn test_model_input_output_ids() {
let buffer = generate_model_buffer(ModelFormat::V2);
let model = Model::load(buffer).unwrap();
let input_id = model.find_node("input").unwrap();
let output_id = model.find_node("output").unwrap();
assert_eq!(model.input_ids(), &[input_id]);
assert_eq!(model.output_ids(), &[output_id]);
assert_eq!(model.node_id("input").ok(), Some(input_id));
assert_eq!(model.find_node("does_not_exist"), None);
let err = model.node_id("does_not_exist").err().unwrap();
assert_eq!(err.node_path(), [Some("does_not_exist")]);
assert_eq!(err.kind(), RunErrorKind::NodeNotFound);
}
#[test]
fn test_unsupported_operator() {
let buffer = generate_model_buffer(ModelFormat::V2);
let registry = op_registry!();
let result = ModelOptions::with_ops(registry).load(buffer);
assert_eq!(
result.err().map(|err| err.to_string()).as_deref(),
Some(
"in node \"concat\": operator error: Concat operator not supported or not enabled"
)
);
}
#[test]
fn test_subset_of_ops_enabled() {
let buffer = generate_model_buffer(ModelFormat::V2);
let registry = op_registry!(Concat, Relu);
let result = ModelOptions::with_ops(registry).load(buffer);
assert!(result.is_ok());
}
#[test]
fn test_shape_info() {
let buffer = generate_model_buffer(ModelFormat::V2);
let model = Model::load(buffer).unwrap();
let input_id = model.input_ids()[0];
let shape = model
.node_info(input_id)
.and_then(|ni| ni.shape())
.expect("input shape missing");
assert_eq!(shape, &[1, 2, 2].map(Dimension::Fixed));
}
#[test]
fn test_value_dtype_info() {
let buffer = generate_model_buffer(ModelFormat::V2);
let model = Model::load(buffer).unwrap();
let input_id = model.input_ids()[0];
let dtype = model
.node_info(input_id)
.and_then(|ni| ni.dtype())
.expect("input dtype missing");
assert_eq!(dtype, ValueType::Tensor(DataType::Float));
}
#[test]
fn test_metadata() {
let buffer = generate_model_buffer(ModelFormat::V2);
let model = Model::load(buffer).unwrap();
assert_eq!(model.metadata().onnx_hash(), Some("abc"));
assert_eq!(model.metadata().description(), None);
}
#[test]
fn test_input_shape() {
let buffer = generate_model_buffer(ModelFormat::V2);
let model = Model::load(buffer).unwrap();
assert_eq!(
model.input_shape(0),
Some(vec![
Dimension::Fixed(1),
Dimension::Fixed(2),
Dimension::Fixed(2),
])
);
}
#[test]
fn test_load_and_run_model() {
struct Case {
format: ModelFormat,
opts: Option<ModelOptions>,
}
let cases = [
Case {
format: ModelFormat::V1,
opts: None,
},
Case {
format: ModelFormat::V2,
opts: None,
},
Case {
format: ModelFormat::V2,
opts: Some({
let mut opts = ModelOptions::with_all_ops();
opts.enable_optimization(false);
opts
}),
},
Case {
format: ModelFormat::V2,
opts: Some({
let mut opts = ModelOptions::with_all_ops();
opts.prepack_weights(true);
opts
}),
},
];
for Case { format, opts } in cases {
let buffer = generate_model_buffer(format);
let model = if let Some(opts) = opts {
opts.load(buffer).unwrap()
} else {
Model::load(buffer).unwrap()
};
let input_id = model.input_ids()[0];
let output_id = model.output_ids()[0];
let input = generate_input();
let result = model
.run(vec![(input_id, input.view().into())], &[output_id], None)
.unwrap();
let result_tensor = check_output(result);
let partial_run_result = model
.partial_run(vec![(input_id, input.into())], &[output_id], None)
.unwrap();
assert_eq!(
partial_run_result,
vec![(output_id, Value::FloatTensor(result_tensor))]
);
}
}
#[test]
fn test_model_debug() {
let buffer = generate_model_buffer(ModelFormat::V2);
let model = Model::load(buffer).unwrap();
let debug_str = format!("{model:?}");
assert_eq!(
debug_str,
"Model { inputs: [\"input\"], outputs: [\"output\"] }"
);
}
#[test]
fn test_load_invalid_model() {
struct Case {
buf: Vec<u8>,
expected_error: &'static str,
}
let buf = generate_model_buffer(ModelFormat::V2);
let mut invalid_model = buf.clone();
let header_size = 32;
invalid_model.insert(header_size, 0);
let mut truncated_buf = buf.clone();
truncated_buf.truncate(truncated_buf.len() - 1);
let cases = [
Case {
buf: b"RTENabc".to_vec(),
expected_error: "invalid header",
},
Case {
buf: invalid_model,
expected_error: "parse error:",
},
Case {
buf: truncated_buf,
expected_error: "graph error: invalid tensor data offset",
},
];
for Case {
buf,
expected_error,
} in cases
{
let err = Model::load(buf).err().unwrap();
assert!(
err.to_string().contains(expected_error),
"expected \"{}\" to contain \"{}\"",
err,
expected_error
);
}
}
#[test]
fn test_load_static_slice() {
let buffer = generate_model_buffer(ModelFormat::V2).leak();
let model = Model::load_static_slice(buffer).unwrap();
let input = generate_input();
let input_id = model.input_ids()[0];
let output_id = model.output_ids()[0];
let result = model
.run(vec![(input_id, input.into())], &[output_id], None)
.unwrap();
check_output(result);
}
#[test]
fn test_load_file() {
let buffer = generate_model_buffer(ModelFormat::V2);
std::fs::write("model-load-file-test.rten", buffer).unwrap();
let model = Model::load_file("model-load-file-test.rten").unwrap();
let input_id = model.input_ids()[0];
let output_id = model.output_ids()[0];
let input = generate_input();
let result = model
.run(vec![(input_id, input.into())], &[output_id], None)
.unwrap();
check_output(result);
}
#[cfg(feature = "mmap")]
#[cfg(not(target_arch = "wasm32"))]
#[test]
fn test_load_mmap() {
let buffer = generate_model_buffer(ModelFormat::V2);
std::fs::write("model-load-mmap-test.rten", buffer).unwrap();
let model = unsafe { Model::load_mmap("model-load-mmap-test.rten").unwrap() };
let input_id = model.input_ids()[0];
let output_id = model.output_ids()[0];
let input = generate_input();
let result = model
.run(vec![(input_id, input.into())], &[output_id], None)
.unwrap();
check_output(result);
}
#[test]
fn test_load_unknown_type() {
let err = Model::load_file("README.md").err().unwrap();
assert_eq!(err.kind(), LoadErrorKind::UnknownFileType);
}
#[cfg(feature = "onnx_format")]
#[test]
fn test_load_onnx() {
let check_model = |model: Model| {
assert_eq!(model.input_ids().len(), 1);
let input_info = model.node_info(model.input_ids()[0]).unwrap();
assert_eq!(input_info.name().unwrap(), "input");
assert_eq!(
input_info.shape().unwrap(),
[1, 1, 28, 28].map(Dimension::Fixed)
);
assert_eq!(model.output_ids().len(), 1);
let output_info = model.node_info(model.output_ids()[0]).unwrap();
assert_eq!(output_info.name().unwrap(), "logits");
assert_eq!(output_info.shape().unwrap(), [1, 10].map(Dimension::Fixed));
let result = model
.run_one(NdTensor::full([1, 1, 28, 28], 0.5).into(), None)
.unwrap();
assert_eq!(result.shape().as_slice(), &[1, 10]);
};
let model_path = "rten-onnx/test-data/mnist.onnx";
let external_model_path = "rten-onnx/test-data/mnist-external/mnist.onnx";
let model = Model::load_file(model_path).unwrap();
check_model(model);
let model = Model::load_file(external_model_path).unwrap();
check_model(model);
let onnx_buf = std::fs::read(model_path).unwrap();
let model = Model::load(onnx_buf).unwrap();
check_model(model);
let onnx_buf = std::fs::read(external_model_path).unwrap();
let data_buf = std::fs::read(format!("{}.data", external_model_path)).unwrap();
let model = ModelOptions::with_all_ops()
.external_data("mnist.onnx.data", data_buf)
.load(onnx_buf)
.unwrap();
check_model(model);
#[cfg(feature = "mmap")]
{
let model = unsafe { Model::load_mmap(external_model_path) }.unwrap();
check_model(model);
}
}
#[test]
fn test_run_one() {
let buffer = generate_model_buffer(ModelFormat::V2);
let model = Model::load(buffer).unwrap();
let input = Tensor::from([[[1., 2.], [-1., -2.]]]);
let result: Tensor<f32> = model
.run_one(input.into(), None)
.unwrap()
.try_into()
.unwrap();
assert_eq!(result.shape(), &[2, 2, 2]);
assert_eq!(result.to_vec(), &[0.5, 0., 0.1, 0., 1., 2., 0., 0.]);
}
#[test]
fn test_omitted_optional_inputs() {
let mut builder = ModelBuilder::new(ModelFormat::V2);
let mut graph_builder = builder.graph_builder();
let output_node = graph_builder.add_value("output", None, None);
graph_builder.add_output(output_node);
graph_builder.add_operator(
"shape",
OpType::Shape(Shape::default()),
&[None],
&[output_node],
);
let graph = graph_builder.finish();
builder.set_graph(graph);
let buffer = builder.finish();
let model = ModelOptions::with_all_ops()
.enable_optimization(false)
.load(buffer)
.unwrap();
let err = model.run(vec![], &[output_node], None).err().unwrap();
assert_eq!(err.node_path(), [Some("shape")]);
assert_eq!(err.kind(), RunErrorKind::OperatorError);
}
#[test]
fn test_all_op_types() {
let mut builder = ModelBuilder::new(ModelFormat::V2);
let mut graph_builder = builder.graph_builder();
let input_node = graph_builder.add_value("input", None, None);
let input_2d = graph_builder.add_value("input.2d", None, None);
let input_bool = graph_builder.add_value("input.bool", None, None);
let input_u8 = graph_builder.add_value("input.u8", None, None);
let input_2d_u8 = graph_builder.add_value("input.2d.u8", None, None);
let input_2d_i8 = graph_builder.add_value("input.2d.i8", None, None);
let input_shape = [1, 1, 3, 3];
let kernel_val = Tensor::from_data(&[1, 1, 1, 1], vec![0.5]);
let kernel = graph_builder.add_constant(kernel_val.view());
let kernel_val_i8 = Tensor::from_data(&[1, 1, 1, 1], vec![0i8]);
let kernel_i8 = graph_builder.add_constant(kernel_val_i8.view());
let mut op_outputs = Vec::new();
let mut add_operator =
|builder: &mut GraphBuilder, name: &str, op: OpType, input_nodes: &[Option<NodeId>]| {
let output_name = format!("{}_out", name);
let op_output_node = builder.add_value(&output_name, None, None);
builder.add_operator(name, op, input_nodes, &[op_output_node]);
op_outputs.push(output_name);
op_output_node
};
macro_rules! add_operator {
($op_name:ident, $op_inputs:expr) => {
add_operator(
&mut graph_builder,
stringify!($op_name),
OpType::$op_name,
&$op_inputs.map(Some),
)
};
($op_name:ident, $op_inputs:expr, $attrs: tt) => {
add_operator(
&mut graph_builder,
stringify!($op_name),
OpType::$op_name(ops::$op_name $attrs),
&$op_inputs.map(Some),
)
};
}
add_operator!(Abs, [input_node]);
add_operator!(Acos, [input_node]);
add_operator!(Acosh, [input_node]);
add_operator!(Add, [input_node, input_node]);
add_operator!(And, [input_bool, input_bool]);
add_operator!(ArgMax, [input_node], { axis: 3, keep_dims: false });
add_operator!(ArgMin, [input_node], { axis: 3, keep_dims: false });
add_operator!(Asin, [input_node]);
add_operator!(Asinh, [input_node]);
add_operator!(Atan, [input_node]);
add_operator!(Atanh, [input_node]);
add_operator!(AveragePool, [input_node], {
kernel_size: [2, 2].into(),
strides: [2, 2].into(),
padding: [0, 0, 0, 0].into(),
count_include_pad: false,
ceil_mode: false,
});
let batch_norm_param_val = Tensor::from([1.0]);
let batch_norm_param = graph_builder.add_constant(batch_norm_param_val.view());
add_operator!(
BatchNormalization,
[
input_node,
batch_norm_param,
batch_norm_param,
batch_norm_param,
batch_norm_param,
],
{ epsilon: 1e-5 }
);
add_operator!(Cast, [input_node], { to: DataType::Float });
add_operator!(CastLike, [input_node, input_node], {});
add_operator!(Ceil, [input_node]);
let clip_min = graph_builder.add_constant(Tensor::from(1.).view());
let clip_max = graph_builder.add_constant(Tensor::from(6.).view());
add_operator!(Clip, [input_node, clip_min, clip_max]);
add_operator!(Concat, [input_node, input_node], { axis: 0 });
let shape = graph_builder.add_constant(Tensor::from([1, 5, 10]).view());
add_operator!(ConstantOfShape, [shape], { value: Scalar::Int32(42) });
add_operator!(Conv, [input_node, kernel], {
dilations: vec![1, 1],
groups: 1,
padding: [1, 1, 1, 1].into(),
strides: vec![1, 1],
});
add_operator!(ConvInteger, [input_u8, kernel_i8], {
dilations: vec![1, 1],
groups: 1,
padding: [1, 1, 1, 1].into(),
strides: vec![1, 1],
});
add_operator!(ConvTranspose, [input_node, kernel], {
strides: vec![2, 2],
padding: [0, 0, 0, 0].into(),
groups: 1,
output_padding: None,
});
add_operator!(Cos, [input_node]);
add_operator!(Cosh, [input_node]);
let const_u8_val = Tensor::from([0u8, 1, 2, 3, 4]);
let const_u8 = graph_builder.add_constant(const_u8_val.view());
let const_f32_val = const_u8_val.map(|x| *x as f32);
let const_f32 = graph_builder.add_constant(const_f32_val.view());
let scale_val = Tensor::from(1.);
let scale = graph_builder.add_constant(scale_val.view());
let zero_point_val = Tensor::from(0u8);
let zero_point = graph_builder.add_constant(zero_point_val.view());
add_operator!(DequantizeLinear, [const_u8, scale, zero_point], {
axis: 0,
});
add_operator!(DepthToSpace, [input_node], {
mode: DepthToSpaceMode::DepthColumnRow,
block_size: 1,
});
add_operator!(QuantizeLinear, [const_f32, scale, zero_point], {
axis: 0,
output_dtype: None,
});
add_operator!(Div, [input_node, input_node]);
#[cfg(feature = "random")]
{
let dropout_out = graph_builder.add_value("Dropout_out", None, None);
let dropout_out_mask = graph_builder.add_value("Dropout_out_mask", None, None);
graph_builder.add_operator(
"Dropout",
OpType::Dropout(ops::Dropout { seed: None }),
&[input_2d].map(Some),
&[dropout_out, dropout_out_mask],
);
}
add_operator!(Elu, [input_node], { alpha: 1.0 });
add_operator!(Equal, [input_node, input_node]);
add_operator!(Erf, [input_node]);
add_operator!(Exp, [input_node]);
let expand_shape_val = Tensor::from([2, 2, 3, 3]);
let expand_shape = graph_builder.add_constant(expand_shape_val.view());
add_operator!(Expand, [input_node, expand_shape]);
add_operator!(EyeLike, [input_2d], { k: 2, dtype: None });
add_operator!(Flatten, [input_node], { axis: 1 });
add_operator!(Floor, [input_node]);
let gather_indices_val = Tensor::from([0]);
let gather_indices = graph_builder.add_constant(gather_indices_val.view());
add_operator!(Gather, [input_node, gather_indices], { axis: 0 });
let gather_elements_indices_val = Tensor::<i32>::zeros(&input_shape);
let gather_elements_indices =
graph_builder.add_constant(gather_elements_indices_val.view());
add_operator!(GatherElements, [input_node, gather_elements_indices], { axis: 0 });
add_operator!(Gelu, [input_node], { approximate: false });
add_operator!(Gemm, [input_2d, input_2d], {
alpha: 1.0,
beta: 1.0,
transpose_a: false,
transpose_b: false,
});
add_operator!(GlobalAveragePool, [input_node]);
add_operator!(GlobalMaxPool, [input_node]);
add_operator!(Greater, [input_node, input_node]);
add_operator!(GreaterOrEqual, [input_node, input_node]);
add_operator!(HardSigmoid, [input_node], {
alpha: 0.2,
beta: 0.5,
});
add_operator!(HardSwish, [input_node]);
add_operator!(Identity, [input_node]);
let if_cond_val = Tensor::from(1);
let if_cond = graph_builder.add_constant(if_cond_val.view());
let mut then_branch_builder = graph_builder.subgraph_builder();
let then_out_val = Tensor::from(2);
let then_out = then_branch_builder.add_constant(then_out_val.view());
then_branch_builder.add_output(then_out);
let then_branch = then_branch_builder.finish();
let mut else_branch_builder = graph_builder.subgraph_builder();
let else_out_val = Tensor::from(3);
let else_out = else_branch_builder.add_constant(else_out_val.view());
else_branch_builder.add_output(else_out);
let else_branch = else_branch_builder.finish();
add_operator(
&mut graph_builder,
"If",
OpType::If(IfArgs {
then_branch,
else_branch,
}),
&[Some(if_cond)],
);
let instance_norm_scale_val = Tensor::from([1.0]);
let instance_norm_scale = graph_builder.add_constant(instance_norm_scale_val.view());
let instance_norm_bias_val = Tensor::from([1.0]);
let instance_norm_bias = graph_builder.add_constant(instance_norm_bias_val.view());
add_operator!(InstanceNormalization, [
input_node, instance_norm_scale, instance_norm_bias
], { epsilon: Some(1e-5) });
add_operator!(IsInf, [input_node]);
add_operator!(IsNaN, [input_node]);
let layer_norm_scale_val = Tensor::full(&[input_shape[input_shape.len() - 1]], 1.);
let layer_norm_scale = graph_builder.add_constant(layer_norm_scale_val.view());
let layer_norm_bias_val = layer_norm_scale_val.clone();
let layer_norm_bias = graph_builder.add_constant(layer_norm_bias_val.view());
add_operator!(LayerNormalization, [
input_node, layer_norm_scale, layer_norm_bias
], { axis: -1, epsilon: Some(1e-5) });
add_operator!(LeakyRelu, [input_node], { alpha: 0.01 });
add_operator!(Less, [input_node, input_node]);
add_operator!(LessOrEqual, [input_node, input_node]);
add_operator!(Log, [input_node]);
add_operator!(LogSoftmax, [input_node], { axis: 1 });
add_operator!(MatMul, [input_2d, input_2d]);
add_operator!(MatMulInteger, [input_2d_u8, input_2d_i8]);
add_operator!(Max, [input_node, input_node]);
add_operator!(MaxPool, [input_node], {
kernel_size: [2, 2].into(),
strides: [2, 2].into(),
padding: [0, 0, 0, 0].into(),
ceil_mode: false,
});
add_operator!(Mean, [input_node, input_node]);
add_operator!(Min, [input_node, input_node]);
add_operator!(Mod, [input_node, input_node], {
fmod: false,
});
add_operator!(Mul, [input_node, input_node]);
add_operator!(Neg, [input_node]);
let nms_n_boxes = 10;
let nms_n_classes = 20;
let nms_boxes =
graph_builder.add_constant(Tensor::<f32>::zeros(&[1, nms_n_boxes, 4]).view());
let nms_scores = graph_builder
.add_constant(Tensor::<f32>::zeros(&[1, nms_n_classes, nms_n_boxes]).view());
let nms_max_outputs_per_class = graph_builder.add_constant(Tensor::from(10).view());
let nms_iou_threshold = graph_builder.add_constant(Tensor::from(0.45).view());
let nms_score_threshold = graph_builder.add_constant(Tensor::from(0.2).view());
add_operator!(NonMaxSuppression, [nms_boxes, nms_scores, nms_max_outputs_per_class, nms_iou_threshold, nms_score_threshold], {
box_order: BoxOrder::CenterWidthHeight,
});
add_operator!(NonZero, [input_node]);
add_operator!(Not, [input_bool]);
let onehot_indices = graph_builder.add_constant(Tensor::from([0, 1, 2]).view());
let onehot_depth = graph_builder.add_constant(Tensor::from(5).view());
let onehot_values = graph_builder.add_constant(Tensor::from([1., 0.]).view());
add_operator!(OneHot, [onehot_indices, onehot_depth, onehot_values], {
axis: -1,
});
add_operator!(Or, [input_bool, input_bool]);
let pads = graph_builder.add_constant(Tensor::from([0, 0, 1, 1, 0, 0, 1, 1]).view());
add_operator!(Pad, [input_node, pads]);
add_operator!(Pow, [input_node, input_node]);
#[cfg(feature = "random")]
{
add_operator!(RandomNormal, [], {
shape: vec![50, 50],
mean: 0.,
scale: 1.,
seed: None,
});
add_operator!(RandomNormalLike, [input_node], {
mean: 0.,
scale: 1.,
seed: None,
});
add_operator!(RandomUniform, [], {
shape: vec![50, 50],
low: 0.,
high: 1.,
seed: None,
});
add_operator!(RandomUniformLike, [input_node], {
low: 0.,
high: 1.,
seed: None,
});
add_operator!(Multinomial, [input_2d], {
sample_size: 4,
seed: None,
});
}
let range_start_node = graph_builder.add_value("range_start", None, None);
let range_limit_node = graph_builder.add_value("range_limit", None, None);
let range_delta_node = graph_builder.add_value("range_delta", None, None);
let range_out = add_operator!(
Range,
[range_start_node, range_limit_node, range_delta_node]
);
add_operator!(Reciprocal, [input_node]);
add_operator!(ReduceMean, [input_node], {
axes: None,
keep_dims: false,
noop_with_empty_axes: false,
});
add_operator!(ReduceMax, [input_node], {
axes: None,
keep_dims: false,
noop_with_empty_axes: false,
});
add_operator!(ReduceMin, [input_node], {
axes: None,
keep_dims: false,
noop_with_empty_axes: false,
});
add_operator!(ReduceProd, [input_node], {
axes: None,
keep_dims: false,
noop_with_empty_axes: false,
});
add_operator!(ReduceSum, [input_node], {
axes: None,
keep_dims: false,
noop_with_empty_axes: false,
});
add_operator!(ReduceSumSquare, [input_node], {
axes: None,
keep_dims: false,
noop_with_empty_axes: false,
});
add_operator!(ReduceL1, [input_node], {
axes: None,
keep_dims: false,
noop_with_empty_axes: false,
});
add_operator!(ReduceL2, [input_node], {
axes: None,
keep_dims: false,
noop_with_empty_axes: false,
});
add_operator!(Relu, [input_node]);
let new_shape = graph_builder.add_constant(Tensor::from([9]).view());
add_operator!(Reshape, [input_node, new_shape], {
allow_zero: false,
});
let resize_roi_val = Tensor::from([0., 0., 0., 0., 1., 1., 1., 1.]);
let resize_scales_val = Tensor::from([1., 1., 2., 2.]);
let resize_roi = graph_builder.add_constant(resize_roi_val.view());
let resize_scales = graph_builder.add_constant(resize_scales_val.view());
add_operator!(Resize, [input_node, resize_roi, resize_scales], {
mode: ResizeMode::Nearest,
nearest_mode: NearestMode::default(),
coord_mode: CoordTransformMode::default()
});
add_operator!(Round, [input_node]);
let upsample_scales = graph_builder.add_constant(Tensor::from([1., 1., 2., 2.]).view());
add_operator!(Upsample, [input_node, upsample_scales], {
mode: ResizeMode::Nearest
});
add_operator!(Shape, [input_node], {
start: Some(1),
end: Some(-1),
});
add_operator!(Sigmoid, [input_node]);
add_operator!(Sign, [input_node]);
add_operator!(Sin, [input_node]);
add_operator!(Sinh, [input_node]);
add_operator!(Size, [input_node]);
let scatter_elem_indices_val = Tensor::<i32>::zeros(&input_shape);
let scatter_elem_indices = graph_builder.add_constant(scatter_elem_indices_val.view());
let scatter_elem_updates_val = Tensor::<f32>::zeros(&input_shape);
let scatter_elem_updates = graph_builder.add_constant(scatter_elem_updates_val.view());
add_operator!(
ScatterElements,
[input_node, scatter_elem_indices, scatter_elem_updates],
{ axis: 0, reduction: None }
);
add_operator!(
Scatter,
[input_node, scatter_elem_indices, scatter_elem_updates],
{ axis: 0 }
);
let rotary_cos = graph_builder.add_constant(Tensor::<f32>::zeros(&[3, 1]).view());
let rotary_sin = graph_builder.add_constant(Tensor::<f32>::zeros(&[3, 1]).view());
let rotary_pos = graph_builder.add_constant(Tensor::from([[0i32, 1, 2]]).view());
add_operator!(
RotaryEmbedding,
[input_node, rotary_cos, rotary_sin, rotary_pos],
{ interleaved: false, num_heads: 1, rotary_embedding_dim: 2 }
);
let const_0 = graph_builder.add_constant(Tensor::from([0]).view());
let const_1 = graph_builder.add_constant(Tensor::from([1]).view());
add_operator!(Slice, [input_node, const_0, const_1, const_0]);
add_operator!(Softplus, [input_node]);
add_operator!(Softmax, [input_node], { axis: 1, flush_nans_to_zero: false });
add_operator!(Sqrt, [input_node]);
add_operator!(Squeeze, [input_node]);
let split_splits = graph_builder.add_constant(Tensor::from([1, 2]).view());
let split_out_1 = graph_builder.add_value("Split_out_1", None, None);
let split_out_2 = graph_builder.add_value("Split_out_2", None, None);
graph_builder.add_operator(
"Split",
OpType::Split(ops::Split {
axis: 1,
num_outputs: None,
}),
&[input_2d, split_splits].map(Some),
&[split_out_1, split_out_2],
);
add_operator!(Sub, [input_node, input_node]);
add_operator!(Sum, [input_node, input_node]);
add_operator!(Tan, [input_node]);
add_operator!(Tanh, [input_node]);
let tile_repeats = graph_builder.add_constant(Tensor::from([1, 2, 3, 4]).view());
add_operator!(Tile, [input_node, tile_repeats]);
let topk_k = graph_builder.add_constant(Tensor::from(3).view());
let topk_out_values = graph_builder.add_value("TopK_out_values", None, None);
let topk_out_indices = graph_builder.add_value("TopK_out_indices", None, None);
graph_builder.add_operator(
"TopK",
OpType::TopK(ops::TopK {
largest: true,
sorted: true,
axis: Some(-1),
}),
&[input_2d, topk_k].map(Some),
&[topk_out_values, topk_out_indices],
);
add_operator!(Transpose, [input_node], { perm: None });
add_operator!(Trilu, [input_node], { upper: true });
let unsqueeze_axes = graph_builder.add_constant(Tensor::from([0, 4]).view());
add_operator!(Unsqueeze, [input_node, unsqueeze_axes]);
let where_cond = graph_builder.add_value("where_cond", None, None);
let where_x = graph_builder.add_value("where_x", None, None);
let where_y = graph_builder.add_value("where_y", None, None);
let where_out = add_operator!(Where, [where_cond, where_x, where_y]);
add_operator!(Xor, [input_bool, input_bool]);
let graph = graph_builder.finish();
builder.set_graph(graph);
let buffer = builder.finish();
let model = Model::load(buffer).unwrap();
let input = Tensor::from_data(&input_shape, vec![1., 2., 3., 4., 5., 6., 7., 8., 9.]);
let input_2d_data = NdTensor::from([[1, 2, 3], [4, 5, 6]]);
let input_bool_data: Tensor<i32> = Tensor::from([0, 1, 1]);
let input_u8_data = input.map(|&x| x as u8);
let input_2d_u8_data = Tensor::from([[1u8, 2], [3, 4]]);
let input_2d_i8_data = Tensor::from([[1i8, 2], [3, 4]]);
for output in op_outputs {
if [
"Dropout_out",
"Dropout_out_mask",
"Gemm_out",
"MatMul_out",
"Multinomial_out",
"Range_out",
"Split_out_1",
"Split_out_2",
"TopK_out_indices",
"TopK_out_values",
"Where_out",
]
.contains(&output.as_str())
{
continue;
}
let output_id = model.find_node(&output).unwrap();
let result = model
.run(
vec![
(input_node, input.view().into()),
(input_bool, input_bool_data.view().into()),
(input_u8, input_u8_data.view().into()),
(input_2d, input_2d_data.view().into()),
(input_2d_u8, input_2d_u8_data.view().into()),
(input_2d_i8, input_2d_i8_data.view().into()),
],
&[output_id],
None,
)
.unwrap();
assert_eq!(result.len(), 1);
let output_id = model.find_node(&output).unwrap();
let result = model
.run(
vec![
(input_node, input.clone().into()),
(input_bool, input_bool_data.clone().into()),
(input_u8, input_u8_data.clone().into()),
(input_2d, input_2d_data.clone().into()),
(input_2d_u8, input_2d_u8_data.view().into()),
(input_2d_i8, input_2d_i8_data.view().into()),
],
&[output_id],
None,
)
.unwrap();
assert_eq!(result.len(), 1);
}
#[allow(unused_mut)]
let mut outputs = vec![
"Gemm_out",
"MatMul_out",
"Split_out_1",
"Split_out_2",
"TopK_out_indices",
"TopK_out_values",
];
#[cfg(feature = "random")]
{
outputs.extend(["Dropout_out", "Dropout_out_mask", "Multinomial_out"]);
}
let input = Tensor::from_data(&[3, 3], vec![1., 2., 3., 4., 5., 6., 7., 8., 9.]);
for output in outputs {
let output_id = model.find_node(output).unwrap();
let result = model
.run(vec![(input_2d, input.view().into())], &[output_id], None)
.unwrap();
assert_eq!(result.len(), 1);
}
let start = Tensor::from(0.);
let limit = Tensor::from(5.);
let delta = Tensor::from(1.);
let result = model
.run(
vec![
(range_start_node, start.into()),
(range_limit_node, limit.into()),
(range_delta_node, delta.into()),
],
&[range_out],
None,
)
.unwrap();
assert_eq!(result.len(), 1);
let cond = Tensor::from(1);
let x = Tensor::from([1, 2, 3]);
let y = Tensor::from([4, 5, 6]);
let result = model
.run(
vec![
(where_cond, cond.into()),
(where_x, x.into()),
(where_y, y.into()),
],
&[where_out],
None,
)
.unwrap();
assert_eq!(result.len(), 1);
}
}