use serde_onnx::export::{GraphBuilder, ToOnnx, ValueRef, export_model};
use serde_onnx::import::{DecodedModel, PayloadVisitor};
use serde_onnx::ir::{Dim, ElemType, Node, ValueType};
use serde_onnx::ml::{ClassLabels, LinearClassifier, NodePayload};
use serde_onnx::proto::encode_model;
struct TinyClassifier {
op: LinearClassifier,
}
impl TinyClassifier {
fn new() -> Self {
TinyClassifier {
op: LinearClassifier {
classlabels: ClassLabels {
ints: Some(vec![0, 1]),
strings: None,
},
coefficients: vec![
0.5, -0.3, -0.5, 0.8, ],
intercepts: Some(vec![0.1, -0.1]),
multi_class: None,
post_transform: None,
},
}
}
}
impl ToOnnx for TinyClassifier {
fn to_graph(
&self,
builder: &mut GraphBuilder,
) -> Result<ValueRef, serde_onnx::export::ExportError> {
let input_ty = ValueType::tensor(ElemType::Float, Some(vec![Dim::Unknown, Dim::Fixed(2)]));
let x = builder.input("X", input_ty)?;
let mut outs = builder.emit_op(&self.op, vec![x], 2)?;
let z = outs.pop().expect("second output (scores)");
let y = outs.pop().expect("first output (labels)");
builder.output(
y.name().to_string(),
ValueType::tensor(ElemType::Int64, Some(vec![Dim::Unknown])),
)?;
builder.output(
z.name().to_string(),
ValueType::tensor(ElemType::Float, Some(vec![Dim::Unknown, Dim::Fixed(2)])),
)?;
Ok(y)
}
}
#[derive(Default)]
struct Describer {
typed: usize,
raw: usize,
}
impl PayloadVisitor for Describer {
fn visit_typed(&mut self, index: usize, payload: &NodePayload) {
self.typed += 1;
let node = payload.node();
println!(
" [{index}] TYPED {} :: {:?} (recognized)",
node.op_type, node.domain
);
}
fn visit_raw(&mut self, index: usize, node: &Node) {
self.raw += 1;
println!(
" [{index}] RAW {} :: {:?} (preserved untyped)",
node.op_type, node.domain
);
}
}
fn main() {
let model = export_model(&TinyClassifier::new(), "tiny_classifier").expect("export");
println!("export: {} node(s)", model.graph.nodes.len());
let bytes = encode_model(&model);
let path = std::env::temp_dir().join("serde_onnx_linear_classifier.onnx");
std::fs::write(&path, &bytes).expect("write .onnx");
println!("saved to {} ({} bytes)", path.display(), bytes.len());
let decoded = DecodedModel::decode_file(&path).expect("decode_file");
println!("import: graph {:?}, opsets:", decoded.graph_name());
for opset in decoded.opset_import() {
println!(" {:?} v{}", opset.domain, opset.version);
}
println!(
"inputs: {:?}",
decoded.inputs().iter().map(|i| &i.name).collect::<Vec<_>>()
);
println!(
"outputs: {:?}",
decoded
.outputs()
.iter()
.map(|o| &o.name)
.collect::<Vec<_>>()
);
let mut describer = Describer::default();
decoded.visit(&mut describer);
assert_eq!(describer.typed, 1, "expected 1 typed node");
assert_eq!(describer.raw, 0, "expected no Raw nodes");
assert!(decoded.warnings.is_empty(), "expected no warnings");
match &decoded.payloads[0] {
NodePayload::LinearClassifier(typed) => {
println!("coefficients: {:?}", typed.op.coefficients);
println!("labels: {:?}", typed.op.classlabels.ints);
}
other => panic!("expected LinearClassifier, got {}", other.node().op_type),
}
println!("export → file → import cycle: OK");
}