onnx-export-rs 0.1.1

Export canonical Rust machine-learning models to ONNX
Documentation
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
}

/// Exports an ONNX-ML SVM regressor or one-class SVM.
#[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))
}

/// Exports an integer-labelled ONNX-ML SVM classifier.
///
/// # Errors
///
/// Returns an error when the canonical raw ONNX coefficient layout is invalid.
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,
    )))
}