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}")))
}
fn to_onnx_gpu(&self) -> Result<ModelProto> {
let mut proto = self.to_onnx()?;
native::tensorize_tree_ensembles(&mut proto)?;
Ok(proto)
}
fn export_onnx_gpu(&self, path: impl AsRef<Path>) -> Result<()> {
let proto = self.to_onnx_gpu()?;
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(())
}
#[cfg(feature = "ensemble")]
#[derive(Clone, Copy)]
pub(crate) enum EnsembleAggregation<'a> {
Mean,
HardVote {
classes: &'a [i64],
weights: &'a [f64],
},
SoftVote {
classes: &'a [i64],
weights: &'a [f64],
},
}
#[cfg(feature = "ensemble")]
fn merged_opset_imports<'a>(
protos: impl IntoIterator<Item = &'a ModelProto>,
minimum_default: i64,
) -> Vec<onnx_export_rs::proto::OperatorSetIdProto> {
use std::collections::BTreeMap;
let mut versions = BTreeMap::<String, i64>::new();
versions.insert(String::new(), minimum_default);
for proto in protos {
for import in &proto.opset_import {
versions
.entry(import.domain.clone())
.and_modify(|version| *version = (*version).max(import.version))
.or_insert(import.version);
}
}
versions
.into_iter()
.map(|(domain, version)| onnx_export_rs::proto::OperatorSetIdProto { domain, version })
.collect()
}
#[cfg(feature = "ensemble")]
fn namespace_graph(
proto: ModelProto,
prefix: &str,
replacement_input: &str,
) -> Result<(
Vec<onnx_export_rs::proto::NodeProto>,
Vec<onnx_export_rs::proto::TensorProto>,
String,
)> {
let graph = proto
.graph
.ok_or_else(|| Error::Backend("ensemble member ONNX model has no graph".into()))?;
let input = graph
.input
.first()
.map(|value| value.name.clone())
.ok_or_else(|| Error::Backend("ensemble member ONNX model has no input".into()))?;
let output = graph
.output
.first()
.map(|value| value.name.clone())
.ok_or_else(|| Error::Backend("ensemble member ONNX model has no output".into()))?;
let rename = |name: &str| {
if name == input {
replacement_input.to_string()
} else {
format!("{prefix}{name}")
}
};
let mut nodes = graph.node;
for node in &mut nodes {
node.input = node.input.iter().map(|name| rename(name)).collect();
node.output = node.output.iter().map(|name| rename(name)).collect();
if !node.name.is_empty() {
node.name = format!("{prefix}{}", node.name);
}
}
let mut initializers = graph.initializer;
for initializer in &mut initializers {
initializer.name = rename(&initializer.name);
}
Ok((nodes, initializers, rename(&output)))
}
#[cfg(feature = "ensemble")]
pub(crate) fn combine_onnx(
protos: Vec<ModelProto>,
aggregation: EnsembleAggregation<'_>,
) -> Result<ModelProto> {
use onnx_export_rs::graph_builder::assemble_model;
use onnx_export_rs::proto::GraphProto;
if protos.is_empty() {
return Err(Error::Backend("cannot export an empty ensemble".into()));
}
let input_info = protos[0]
.graph
.as_ref()
.and_then(|graph| graph.input.first())
.cloned()
.ok_or_else(|| Error::Backend("ensemble member ONNX model has no input".into()))?;
let opset_imports = merged_opset_imports(&protos, 13);
let opset = opset_imports
.iter()
.find(|opset| opset.domain.is_empty())
.map_or(13, |opset| opset.version);
let ir = protos
.iter()
.map(|proto| proto.ir_version)
.max()
.unwrap_or(8);
let mut nodes = Vec::new();
let mut initializers = Vec::new();
let mut outputs = Vec::new();
for (index, proto) in protos.into_iter().enumerate() {
let (mut member_nodes, mut member_initializers, output) =
namespace_graph(proto, &format!("mw_m{index}_"), "mw_input")?;
nodes.append(&mut member_nodes);
initializers.append(&mut member_initializers);
outputs.push(output);
}
let final_output = aggregate_ensemble(&mut nodes, &mut initializers, outputs, aggregation)?;
let mut input = input_info;
input.name = "mw_input".into();
let output_info = make_value_info(
final_output.clone(),
&[Dimension::Symbolic("batch".into()), Dimension::Fixed(1)],
);
let mut model = assemble_model(
GraphProto {
node: nodes,
name: "millwright_ensemble".into(),
initializer: initializers,
doc_string: String::new(),
input: vec![input],
output: vec![output_info],
value_info: vec![],
},
opset,
ir,
);
model.opset_import = opset_imports;
Ok(model)
}
#[cfg(feature = "ensemble")]
fn map_class_index(
nodes: &mut Vec<onnx_export_rs::proto::NodeProto>,
initializers: &mut Vec<onnx_export_rs::proto::TensorProto>,
classes: &[i64],
index: &str,
) -> String {
use ndarray::Array1;
let mut terms = Vec::new();
for (position, class) in classes.iter().enumerate() {
let position_name = format!("mw_position_{position}");
let class_name = format!("mw_label_{position}");
let equal = format!("mw_index_eq_{position}");
let cast = format!("mw_index_cast_{position}");
let term = format!("mw_label_term_{position}");
initializers.push(make_tensor(
&position_name,
&Array1::from(vec![position as f32]).into_dyn(),
));
initializers.push(make_tensor(
&class_name,
&Array1::from(vec![*class as f32]).into_dyn(),
));
nodes.push(make_node(
"Equal",
[index, position_name.as_str()],
[equal.as_str()],
vec![],
));
nodes.push(make_node(
"Cast",
[equal.as_str()],
[cast.as_str()],
vec![int_attribute("to", FLOAT as i64)],
));
nodes.push(make_node(
"Mul",
[cast.as_str(), class_name.as_str()],
[term.as_str()],
vec![],
));
terms.push(term);
}
let mut current = terms[0].clone();
for (i, term) in terms.iter().skip(1).enumerate() {
let output = if i + 2 == terms.len() {
"mw_output".into()
} else {
format!("mw_label_sum{i}")
};
nodes.push(make_node(
"Add",
[current.as_str(), term.as_str()],
[output.as_str()],
vec![],
));
current = output;
}
if terms.len() == 1 {
nodes.push(make_node(
"Identity",
[current.as_str()],
["mw_output"],
vec![],
));
"mw_output".into()
} else {
current
}
}
#[cfg(feature = "ensemble")]
pub(crate) fn stack_onnx(bases: Vec<ModelProto>, meta: ModelProto) -> Result<ModelProto> {
use onnx_export_rs::graph_builder::{assemble_model, make_node};
use onnx_export_rs::proto::GraphProto;
if bases.is_empty() {
return Err(Error::Backend(
"cannot export stacking without base models".into(),
));
}
let mut input = bases[0]
.graph
.as_ref()
.and_then(|graph| graph.input.first())
.cloned()
.ok_or_else(|| Error::Backend("stacking base ONNX model has no input".into()))?;
input.name = "mw_input".into();
let opset_imports = merged_opset_imports(bases.iter().chain(std::iter::once(&meta)), 13);
let opset = opset_imports
.iter()
.find(|opset| opset.domain.is_empty())
.map_or(13, |opset| opset.version);
let ir = bases
.iter()
.chain(std::iter::once(&meta))
.map(|proto| proto.ir_version)
.max()
.unwrap_or(8);
let mut nodes = Vec::new();
let mut initializers = Vec::new();
let mut outputs = Vec::new();
for (index, proto) in bases.into_iter().enumerate() {
let (mut member_nodes, mut member_initializers, output) =
namespace_graph(proto, &format!("mw_b{index}_"), "mw_input")?;
nodes.append(&mut member_nodes);
initializers.append(&mut member_initializers);
outputs.push(output);
}
nodes.push(make_node(
"Concat",
outputs,
["mw_meta_input"],
vec![int_attribute("axis", 1)],
));
let (mut meta_nodes, mut meta_initializers, meta_output) =
namespace_graph(meta, "mw_meta_", "mw_meta_input")?;
nodes.append(&mut meta_nodes);
initializers.append(&mut meta_initializers);
let mut model = assemble_model(
GraphProto {
node: nodes,
name: "millwright_stacking".into(),
initializer: initializers,
doc_string: String::new(),
input: vec![input],
output: vec![make_value_info(
meta_output,
&[Dimension::Symbolic("batch".into()), Dimension::Fixed(1)],
)],
value_info: vec![],
},
opset,
ir,
);
model.opset_import = opset_imports;
Ok(model)
}
#[cfg(feature = "ensemble")]
fn scalar_initializer(name: &str, value: f64) -> onnx_export_rs::proto::TensorProto {
use ndarray::Array1;
make_tensor(name, &Array1::from(vec![value as f32]).into_dyn())
}
#[cfg(feature = "ensemble")]
fn add_chain(
nodes: &mut Vec<onnx_export_rs::proto::NodeProto>,
terms: Vec<String>,
stem: &str,
) -> Result<String> {
let mut iter = terms.into_iter();
let mut current = iter
.next()
.ok_or_else(|| Error::Backend("ensemble aggregation has no terms".into()))?;
for (index, term) in iter.enumerate() {
let output = format!("mw_{stem}_sum{index}");
nodes.push(make_node(
"Add",
[current.as_str(), term.as_str()],
[output.as_str()],
vec![],
));
current = output;
}
Ok(current)
}
#[cfg(feature = "ensemble")]
fn aggregate_ensemble(
nodes: &mut Vec<onnx_export_rs::proto::NodeProto>,
initializers: &mut Vec<onnx_export_rs::proto::TensorProto>,
outputs: Vec<String>,
aggregation: EnsembleAggregation<'_>,
) -> Result<String> {
match aggregation {
EnsembleAggregation::Mean => aggregate_mean(nodes, initializers, outputs),
EnsembleAggregation::HardVote { classes, weights } => {
aggregate_hard_vote(nodes, initializers, &outputs, classes, weights)
}
EnsembleAggregation::SoftVote { classes, weights } => {
aggregate_soft_vote(nodes, initializers, &outputs, classes, weights)
}
}
}
#[cfg(feature = "ensemble")]
fn aggregate_mean(
nodes: &mut Vec<onnx_export_rs::proto::NodeProto>,
initializers: &mut Vec<onnx_export_rs::proto::TensorProto>,
outputs: Vec<String>,
) -> Result<String> {
let count = outputs.len();
let sum = add_chain(nodes, outputs, "mean")?;
if count == 1 {
return Ok(sum);
}
initializers.push(scalar_initializer("mw_divisor", count as f64));
nodes.push(make_node(
"Div",
[sum.as_str(), "mw_divisor"],
["mw_output"],
vec![],
));
Ok("mw_output".into())
}
#[cfg(feature = "ensemble")]
fn aggregate_hard_vote(
nodes: &mut Vec<onnx_export_rs::proto::NodeProto>,
initializers: &mut Vec<onnx_export_rs::proto::TensorProto>,
outputs: &[String],
classes: &[i64],
weights: &[f64],
) -> Result<String> {
if outputs.len() != weights.len() || classes.is_empty() {
return Err(Error::Backend(
"invalid hard-voting ONNX aggregation".into(),
));
}
let mut class_scores = Vec::new();
for (class_index, class) in classes.iter().enumerate() {
let class_name = format!("mw_class_{class_index}");
initializers.push(scalar_initializer(&class_name, *class as f64));
let mut terms = Vec::new();
for (member_index, output) in outputs.iter().enumerate() {
let equal = format!("mw_eq_{class_index}_{member_index}");
let cast = format!("mw_cast_{class_index}_{member_index}");
let weighted = format!("mw_weighted_{class_index}_{member_index}");
let weight = format!("mw_weight_{member_index}");
if class_index == 0 {
initializers.push(scalar_initializer(&weight, weights[member_index]));
}
nodes.push(make_node(
"Equal",
[output.as_str(), class_name.as_str()],
[equal.as_str()],
vec![],
));
nodes.push(make_node(
"Cast",
[equal.as_str()],
[cast.as_str()],
vec![int_attribute("to", FLOAT as i64)],
));
nodes.push(make_node(
"Mul",
[cast.as_str(), weight.as_str()],
[weighted.as_str()],
vec![],
));
terms.push(weighted);
}
class_scores.push(add_chain(nodes, terms, &format!("class{class_index}"))?);
}
nodes.push(make_node(
"Concat",
class_scores,
["mw_scores"],
vec![int_attribute("axis", 1)],
));
append_argmax(nodes, "mw_scores");
Ok(map_class_index(nodes, initializers, classes, "mw_index_f"))
}
#[cfg(feature = "ensemble")]
fn aggregate_soft_vote(
nodes: &mut Vec<onnx_export_rs::proto::NodeProto>,
initializers: &mut Vec<onnx_export_rs::proto::TensorProto>,
outputs: &[String],
classes: &[i64],
weights: &[f64],
) -> Result<String> {
if outputs.len() != weights.len() || classes.is_empty() {
return Err(Error::Backend(
"invalid soft-voting ONNX aggregation".into(),
));
}
let mut terms = Vec::new();
for (index, output) in outputs.iter().enumerate() {
let weight = format!("mw_weight_{index}");
let weighted = format!("mw_weighted_{index}");
initializers.push(scalar_initializer(&weight, weights[index]));
nodes.push(make_node(
"Mul",
[output.as_str(), weight.as_str()],
[weighted.as_str()],
vec![],
));
terms.push(weighted);
}
let scores = add_chain(nodes, terms, "soft")?;
append_argmax(nodes, &scores);
Ok(map_class_index(nodes, initializers, classes, "mw_index_f"))
}
#[cfg(feature = "ensemble")]
fn append_argmax(nodes: &mut Vec<onnx_export_rs::proto::NodeProto>, scores: &str) {
nodes.push(make_node(
"ArgMax",
[scores],
["mw_index"],
vec![int_attribute("axis", 1), int_attribute("keepdims", 1)],
));
nodes.push(make_node(
"Cast",
["mw_index"],
["mw_index_f"],
vec![int_attribute("to", FLOAT as i64)],
));
}
pub(crate) fn append_label_map(proto: &mut ModelProto, labels: &[i64]) -> Result<()> {
if labels.iter().copied().eq(0..labels.len() as i64) {
return Ok(());
}
let graph = proto
.graph
.as_mut()
.ok_or_else(|| Error::Backend("exported model has no graph".into()))?;
let index_out = graph
.output
.first()
.map(|o| o.name.clone())
.ok_or_else(|| Error::Backend("exported model has no output".into()))?;
graph.initializer.push(make_i64_tensor(
"mw_labels",
&[labels.len()],
labels.to_vec(),
));
graph.node.push(make_node(
"Gather",
["mw_labels", index_out.as_str()],
["mw_label"],
vec![int_attribute("axis", 0)],
));
if let Some(o) = graph.output.first_mut() {
o.name = "mw_label".into();
}
Ok(())
}
#[derive(Clone)]
pub struct InferenceModel {
backend: std::sync::Arc<Backend>,
}
enum Backend {
Tract(TractPlan),
Native(native::NativeGraph),
#[cfg(feature = "gpu-inference")]
Ort(std::sync::Mutex<ort::session::Session>),
#[cfg(feature = "gpu-inference")]
OrtPool(Vec<std::sync::Mutex<ort::session::Session>>),
}
#[cfg(feature = "gpu-inference")]
#[derive(Clone, Copy, Debug, Default, PartialEq, Eq)]
pub enum Device {
#[default]
Auto,
Gpu,
Cpu,
}
#[cfg(feature = "gpu-inference")]
const HAS_GPU_PROVIDER: bool =
cfg!(feature = "gpu-cuda") || cfg!(target_os = "windows") || cfg!(target_os = "macos");
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)),
})
}
#[cfg(feature = "gpu-inference")]
pub fn load_on(path: impl AsRef<Path>, device: Device) -> Result<InferenceModel> {
let path = path.as_ref();
match device {
Device::Auto => {
if let Ok(session) = build_session(path, gpu_execution_providers(false, None)) {
return Ok(Self::wrap_ort(session));
}
Ok(Self::wrap_ort(build_session(
path,
cpu_execution_providers(),
)?))
}
Device::Gpu => {
if !HAS_GPU_PROVIDER {
return Err(Error::Backend(
"Device::Gpu requires a GPU execution provider, but none is compiled in \
(build on Windows/macOS, or enable the `gpu-cuda` feature)"
.into(),
));
}
Ok(Self::wrap_ort(build_session(
path,
gpu_execution_providers(true, None),
)?))
}
Device::Cpu => Ok(Self::wrap_ort(build_session(
path,
cpu_execution_providers(),
)?)),
}
}
#[cfg(feature = "gpu-inference")]
pub fn load_multi(path: impl AsRef<Path>, device_ids: &[i32]) -> Result<InferenceModel> {
if !HAS_GPU_PROVIDER {
return Err(Error::Backend(
"load_multi requires a GPU execution provider, but none is compiled in \
(build on Windows/macOS, or enable the `gpu-cuda` feature)"
.into(),
));
}
if device_ids.is_empty() {
return Err(Error::Backend(
"load_multi needs at least one device id".into(),
));
}
let path = path.as_ref();
let mut sessions = Vec::with_capacity(device_ids.len());
for &id in device_ids {
let session = build_session(path, gpu_execution_providers(true, Some(id)))?;
sessions.push(std::sync::Mutex::new(session));
}
Ok(Self {
backend: std::sync::Arc::new(Backend::OrtPool(sessions)),
})
}
#[cfg(feature = "gpu-inference")]
fn wrap_ort(session: ort::session::Session) -> InferenceModel {
Self {
backend: std::sync::Arc::new(Backend::Ort(std::sync::Mutex::new(session))),
}
}
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),
#[cfg(feature = "gpu-inference")]
Backend::Ort(session) => Self::ort_predict(session, frame),
#[cfg(feature = "gpu-inference")]
Backend::OrtPool(sessions) => Self::ort_pool_predict(sessions, frame),
}
}
#[cfg(feature = "gpu-inference")]
fn ort_predict(
session: &std::sync::Mutex<ort::session::Session>,
frame: &Frame,
) -> Result<Vec<f64>> {
use ort::value::Tensor;
let (n, p) = frame.shape();
let data: Vec<f32> = frame.buf().iter().map(|v| *v as f32).collect();
let input = Tensor::from_array(([n, p], data.into_boxed_slice()))
.map_err(|e| Error::Backend(format!("onnxruntime input build failed: {e}")))?;
let mut session = session
.lock()
.map_err(|_| Error::Backend("onnxruntime session mutex poisoned".into()))?;
let outputs = session
.run(ort::inputs![input])
.map_err(|e| Error::Backend(format!("onnxruntime run failed: {e}")))?;
let value = &outputs[0];
if let Ok(view) = value.try_extract_array::<i64>() {
return Ok(view.iter().map(|v| *v as f64).collect());
}
let view = value
.try_extract_array::<f32>()
.map_err(|e| Error::Backend(format!("unexpected onnxruntime 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())
}
}
#[cfg(feature = "gpu-inference")]
fn ort_pool_predict(
sessions: &[std::sync::Mutex<ort::session::Session>],
frame: &Frame,
) -> Result<Vec<f64>> {
if sessions.len() == 1 {
return Self::ort_predict(&sessions[0], frame);
}
let (n, _) = frame.shape();
let chunk = n.div_ceil(sessions.len()).max(1);
let rows = frame.as_rows();
let cols = frame.columns().to_vec();
let subframes: Vec<Frame> = rows
.chunks(chunk)
.map(|rs| Frame::from_rows(rs.to_vec(), cols.clone()))
.collect::<Result<_>>()?;
let mut parts: Vec<Result<Vec<f64>>> = Vec::with_capacity(subframes.len());
std::thread::scope(|scope| {
let handles: Vec<_> = subframes
.iter()
.zip(sessions.iter())
.map(|(sf, sess)| scope.spawn(move || Self::ort_predict(sess, sf)))
.collect();
for handle in handles {
parts.push(handle.join().unwrap_or_else(|_| {
Err(Error::Backend("onnxruntime worker thread panicked".into()))
}));
}
});
let mut out = Vec::with_capacity(n);
for part in parts {
out.extend(part?);
}
Ok(out)
}
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())
}
}
}
#[cfg(feature = "gpu-inference")]
fn build_session(
path: &Path,
providers: Vec<ort::ep::ExecutionProviderDispatch>,
) -> Result<ort::session::Session> {
use ort::session::Session;
let mut builder = Session::builder()
.map_err(|e| Error::Backend(format!("onnxruntime builder failed: {e}")))?;
builder = builder
.with_memory_pattern(false)
.map_err(|e| Error::Backend(format!("onnxruntime session option failed: {e}")))?;
builder = builder
.with_execution_providers(providers)
.map_err(|e| Error::Backend(format!("onnxruntime provider setup failed: {e}")))?;
builder
.commit_from_file(path)
.map_err(|e| Error::Backend(format!("onnxruntime load failed: {e}")))
}
#[cfg(feature = "gpu-inference")]
#[allow(clippy::vec_init_then_push)]
fn gpu_execution_providers(
strict: bool,
device_id: Option<i32>,
) -> Vec<ort::ep::ExecutionProviderDispatch> {
let mut providers = Vec::new();
#[cfg(any(feature = "gpu-cuda", target_os = "windows", target_os = "macos"))]
{
let gpu = |dispatch: ort::ep::ExecutionProviderDispatch| {
if strict {
dispatch.error_on_failure()
} else {
dispatch
}
};
#[cfg(feature = "gpu-cuda")]
{
let mut ep = ort::ep::CUDA::default();
if let Some(id) = device_id {
ep = ep.with_device_id(id);
}
providers.push(gpu(ep.build()));
}
#[cfg(target_os = "windows")]
{
let mut ep = ort::ep::DirectML::default();
if let Some(id) = device_id {
ep = ep.with_device_id(id);
}
providers.push(gpu(ep.build()));
}
#[cfg(target_os = "macos")]
{
let _ = device_id; providers.push(gpu(ort::ep::CoreML::default().build()));
}
}
#[cfg(not(any(feature = "gpu-cuda", target_os = "windows", target_os = "macos")))]
let _ = (strict, device_id);
providers.push(ort::ep::CPU::default().build());
providers
}
#[cfg(feature = "gpu-inference")]
fn cpu_execution_providers() -> Vec<ort::ep::ExecutionProviderDispatch> {
vec![ort::ep::CPU::default().build()]
}
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;