use crate::canonical::{SvmClassifier, SvmKernel, SvmRegressor};
use crate::graph_builder::{
assemble_model, floats_attribute, int_attribute, ints_attribute, make_node,
make_typed_value_info, make_value_info, string_attribute, Dimension, INT64,
};
use crate::proto::{GraphProto, ModelProto, OperatorSetIdProto};
use crate::{Result, IR_VERSION, ML_OPSET_VERSION, OPSET_VERSION};
fn kernel_attributes(kernel: SvmKernel) -> Vec<crate::proto::AttributeProto> {
vec![
string_attribute("kernel_type", kernel.onnx_name()),
floats_attribute("kernel_params", kernel.onnx_parameters()),
]
}
fn add_ml_opset(mut model: ModelProto) -> ModelProto {
model.opset_import.push(OperatorSetIdProto {
domain: "ai.onnx.ml".into(),
version: ML_OPSET_VERSION,
});
model
}
#[must_use]
pub fn export_svm_regressor(model: &SvmRegressor) -> ModelProto {
let mut attributes = kernel_attributes(model.kernel);
attributes.extend([
floats_attribute(
"support_vectors",
model
.support_vectors
.iter()
.map(|&value| value as f32)
.collect(),
),
floats_attribute(
"coefficients",
model
.coefficients
.iter()
.map(|&value| value as f32)
.collect(),
),
floats_attribute("rho", vec![model.rho as f32]),
int_attribute("n_supports", model.support_vectors.nrows() as i64),
int_attribute("one_class", i64::from(model.one_class)),
string_attribute("post_transform", b"NONE".as_slice()),
]);
let mut node = make_node("SVMRegressor", ["X"], ["Y"], attributes);
node.domain = "ai.onnx.ml".into();
let graph = GraphProto {
node: vec![node],
name: "svm_regressor".into(),
initializer: Vec::new(),
doc_string: String::new(),
input: vec![make_value_info(
"X",
&[
Dimension::Symbolic("batch".into()),
Dimension::Fixed(model.support_vectors.ncols()),
],
)],
output: vec![make_value_info(
"Y",
&[Dimension::Symbolic("batch".into()), Dimension::Fixed(1)],
)],
value_info: Vec::new(),
};
add_ml_opset(assemble_model(graph, OPSET_VERSION, IR_VERSION))
}
pub fn export_svm_classifier(model: &SvmClassifier) -> Result<ModelProto> {
model.validate()?;
let mut attributes = kernel_attributes(model.kernel);
attributes.extend([
ints_attribute("classlabels_ints", model.class_labels.clone()),
ints_attribute(
"vectors_per_class",
model
.vectors_per_class
.iter()
.map(|&value| value as i64)
.collect(),
),
floats_attribute(
"support_vectors",
model
.support_vectors
.iter()
.map(|&value| value as f32)
.collect(),
),
floats_attribute(
"coefficients",
model
.coefficients
.iter()
.map(|&value| value as f32)
.collect(),
),
floats_attribute("rho", model.rho.iter().map(|&value| value as f32).collect()),
floats_attribute(
"prob_a",
model.prob_a.iter().map(|&value| value as f32).collect(),
),
floats_attribute(
"prob_b",
model.prob_b.iter().map(|&value| value as f32).collect(),
),
string_attribute("post_transform", b"NONE".as_slice()),
]);
let mut node = make_node("SVMClassifier", ["X"], ["label", "scores"], attributes);
node.domain = "ai.onnx.ml".into();
let graph = GraphProto {
node: vec![node],
name: "svm_classifier".into(),
initializer: Vec::new(),
doc_string: String::new(),
input: vec![make_value_info(
"X",
&[
Dimension::Symbolic("batch".into()),
Dimension::Fixed(model.support_vectors.ncols()),
],
)],
output: vec![
make_typed_value_info("label", &[Dimension::Symbolic("batch".into())], INT64),
make_value_info(
"scores",
&[
Dimension::Symbolic("batch".into()),
Dimension::Fixed(model.class_labels.len()),
],
),
],
value_info: Vec::new(),
};
Ok(add_ml_opset(assemble_model(
graph,
OPSET_VERSION,
IR_VERSION,
)))
}