use std::path::Path;
use ndarray::Array2;
use onnx_export_rs::graph_builder::{
int_attribute, make_i64_tensor, make_node, make_tensor, make_value_info, save_to_file,
Dimension, FLOAT,
};
use onnx_export_rs::proto::ModelProto;
use crate::error::{Error, Result};
use crate::frame::Frame;
pub trait ExportOnnx {
fn to_onnx(&self) -> Result<ModelProto>;
fn export_onnx(&self, path: impl AsRef<Path>) -> Result<()> {
let proto = self.to_onnx()?;
save_to_file(&proto, path).map_err(|e| Error::Backend(format!("ONNX save failed: {e}")))
}
}
pub enum Prefix {
Affine { shift: Vec<f64>, scale: Vec<f64> },
Impute { fill: Vec<f64> },
OneHot { columns: Vec<Vec<i64>> },
}
fn row_tensor(name: &str, vals: &[f64]) -> Result<onnx_export_rs::proto::TensorProto> {
let f: Vec<f32> = vals.iter().map(|v| *v as f32).collect();
Ok(make_tensor(
name,
&Array2::from_shape_vec((1, f.len()), f)
.map_err(|e| Error::Backend(e.to_string()))?
.into_dyn(),
))
}
pub(crate) fn prepend_prefixes(proto: &mut ModelProto, prefixes: &[Prefix]) -> Result<()> {
let graph = proto
.graph
.as_mut()
.ok_or_else(|| Error::Backend("exported model has no graph".into()))?;
let est_input = graph
.input
.first()
.map(|vi| vi.name.clone())
.ok_or_else(|| Error::Backend("exported model has no input".into()))?;
let mut nodes = Vec::new();
let mut inits = Vec::new();
let mut cur = "mw_input".to_string();
for (i, prefix) in prefixes.iter().enumerate() {
let out = if i + 1 == prefixes.len() {
est_input.clone()
} else {
format!("mw_pre{i}")
};
match prefix {
Prefix::Impute { fill } => {
let mask = format!("mw_isnan{i}");
let fill_name = format!("mw_fill{i}");
nodes.push(make_node(
"IsNaN",
[cur.as_str()],
[mask.as_str()],
Vec::new(),
));
nodes.push(make_node(
"Where",
[mask.as_str(), fill_name.as_str(), cur.as_str()],
[out.as_str()],
Vec::new(),
));
inits.push(row_tensor(&fill_name, fill)?);
}
Prefix::Affine { shift, scale } => {
let centered = format!("mw_cent{i}");
let shift_name = format!("mw_shift{i}");
let scale_name = format!("mw_scale{i}");
nodes.push(make_node(
"Sub",
[cur.as_str(), shift_name.as_str()],
[centered.as_str()],
Vec::new(),
));
nodes.push(make_node(
"Div",
[centered.as_str(), scale_name.as_str()],
[out.as_str()],
Vec::new(),
));
inits.push(row_tensor(&shift_name, shift)?);
inits.push(row_tensor(&scale_name, scale)?);
}
Prefix::OneHot { columns } => {
let mut pieces: Vec<String> = Vec::new();
for (c, cats) in columns.iter().enumerate() {
let idx_name = format!("mw_idx{i}_{c}");
inits.push(make_i64_tensor(&idx_name, &[1], vec![c as i64]));
let col = format!("mw_col{i}_{c}");
nodes.push(make_node(
"Gather",
[cur.as_str(), idx_name.as_str()],
[col.as_str()],
vec![int_attribute("axis", 1)],
));
if cats.is_empty() {
pieces.push(col);
continue;
}
let rounded = format!("mw_round{i}_{c}");
nodes.push(make_node(
"Round",
[col.as_str()],
[rounded.as_str()],
Vec::new(),
));
for (j, cat) in cats.iter().enumerate() {
let cat_name = format!("mw_cat{i}_{c}_{j}");
inits.push(row_tensor(&cat_name, &[*cat as f64])?);
let eq = format!("mw_eq{i}_{c}_{j}");
nodes.push(make_node(
"Equal",
[rounded.as_str(), cat_name.as_str()],
[eq.as_str()],
Vec::new(),
));
let ind = format!("mw_ind{i}_{c}_{j}");
nodes.push(make_node(
"Cast",
[eq.as_str()],
[ind.as_str()],
vec![int_attribute("to", FLOAT as i64)],
));
pieces.push(ind);
}
}
nodes.push(make_node(
"Concat",
pieces,
[out.as_str()],
vec![int_attribute("axis", 1)],
));
}
}
cur = out;
}
for init in inits {
graph.initializer.push(init);
}
nodes.append(&mut graph.node);
graph.node = nodes;
let raw_width = match &prefixes[0] {
Prefix::Affine { shift, .. } => shift.len(),
Prefix::Impute { fill } => fill.len(),
Prefix::OneHot { columns } => columns.len(),
};
if let Some(vi) = graph.input.first_mut() {
*vi = make_value_info(
"mw_input",
&[
Dimension::Symbolic("batch".into()),
Dimension::Fixed(raw_width),
],
);
}
Ok(())
}
#[derive(Clone)]
pub struct InferenceModel {
backend: std::sync::Arc<Backend>,
}
enum Backend {
Tract(TractPlan),
Native(native::NativeGraph),
}
type TractPlan = std::sync::Arc<tract_onnx::prelude::TypedRunnableModel>;
impl InferenceModel {
pub fn load(path: impl AsRef<Path>) -> Result<InferenceModel> {
let path = path.as_ref();
let bytes =
std::fs::read(path).map_err(|e| Error::Backend(format!("ONNX read failed: {e}")))?;
use prost::Message;
let proto = ModelProto::decode(&bytes[..])
.map_err(|e| Error::Backend(format!("ONNX decode failed: {e}")))?;
if native::needs_native(&proto) {
let graph = native::NativeGraph::from_proto(&proto)?;
return Ok(Self {
backend: std::sync::Arc::new(Backend::Native(graph)),
});
}
use tract_onnx::prelude::*;
let plan = tract_onnx::onnx()
.model_for_path(path)
.map_err(|e| Error::Backend(format!("ONNX load failed: {e}")))?
.into_optimized()
.map_err(|e| Error::Backend(format!("ONNX optimize failed: {e}")))?
.into_runnable()
.map_err(|e| Error::Backend(format!("ONNX plan failed: {e}")))?;
Ok(Self {
backend: std::sync::Arc::new(Backend::Tract(plan)),
})
}
pub fn predict(&self, frame: &Frame) -> Result<Vec<f64>> {
match &*self.backend {
Backend::Native(g) => g.run(frame),
Backend::Tract(plan) => Self::tract_predict(plan, frame),
}
}
fn tract_predict(plan: &TractPlan, frame: &Frame) -> Result<Vec<f64>> {
use tract_onnx::prelude::*;
let (n, p) = frame.shape();
let data: Vec<f32> = frame.buf().iter().map(|v| *v as f32).collect();
let input = tract_ndarray::Array2::from_shape_vec((n, p), data)
.map_err(|e| Error::Backend(e.to_string()))?;
let tensor: Tensor = input.into();
let outputs = plan
.run(tvec!(tensor.into()))
.map_err(|e| Error::Backend(format!("ONNX run failed: {e}")))?;
let out: &Tensor = &outputs[0];
let plain = out
.try_as_plain()
.map_err(|e| Error::Backend(format!("ONNX output not plain: {e}")))?;
if let Ok(view) = plain.to_array_view::<i64>() {
return Ok(view.iter().map(|v| *v as f64).collect());
}
let view = plain
.to_array_view::<f32>()
.map_err(|e| Error::Backend(format!("unexpected ONNX output type: {e}")))?;
let shape = view.shape();
if shape.len() == 2 && shape[1] > 1 {
let cols = shape[1];
let flat: Vec<f32> = view.iter().copied().collect();
Ok((0..n)
.map(|r| {
let row = &flat[r * cols..(r + 1) * cols];
let mut best = 0usize;
for c in 1..cols {
if row[c] > row[best] {
best = c;
}
}
best as f64
})
.collect())
} else {
Ok(view.iter().map(|v| *v as f64).collect())
}
}
}
impl crate::traits::Estimator for InferenceModel {
fn name(&self) -> &'static str {
"InferenceModel"
}
fn fit(&mut self, _dataset: &crate::frame::Dataset) -> Result<()> {
Ok(())
}
}
impl crate::traits::Predictor for InferenceModel {
fn predict(&self, frame: &Frame) -> Result<Vec<f64>> {
InferenceModel::predict(self, frame)
}
}
#[path = "onnx_native.rs"]
mod native;
#[cfg(all(test, feature = "smartcore-backend"))]
#[path = "onnx_tests.rs"]
mod tests;