onnx-export-rs 0.1.1

Export canonical Rust machine-learning models to ONNX
Documentation
use ndarray::Array1;

use crate::canonical::LinearScoreClassifier;
use crate::graph_builder::{
    assemble_model, int_attribute, make_i64_tensor, make_node, make_tensor, make_typed_value_info,
    make_value_info, Dimension, INT64,
};
use crate::proto::{GraphProto, ModelProto};
use crate::{IR_VERSION, OPSET_VERSION};

/// Exports an affine score classifier followed by label lookup.
#[must_use]
pub fn export_linear_score_classifier(model: &LinearScoreClassifier) -> ModelProto {
    let features = model.coefficients.nrows();
    let classes = model.coefficients.ncols();
    let coefficients = model.coefficients.mapv(|value| value as f32);
    let bias = model.bias.mapv(|value| value as f32);
    let mut initializer = vec![
        make_tensor("coefficients", &coefficients.into_dyn()),
        make_tensor("bias", &bias.into_dyn()),
        make_i64_tensor("labels", &[classes], model.class_labels.clone()),
    ];
    let mut nodes = Vec::new();
    let score_input = if let Some(threshold) = model.binarize {
        initializer.push(make_tensor(
            "threshold",
            &Array1::from_vec(vec![threshold as f32]).into_dyn(),
        ));
        nodes.extend([
            make_node("Greater", ["X", "threshold"], ["binary_bool"], Vec::new()),
            make_node(
                "Cast",
                ["binary_bool"],
                ["binary"],
                vec![int_attribute("to", 1)],
            ),
        ]);
        "binary"
    } else {
        "X"
    };
    nodes.extend([
        make_node(
            "Gemm",
            [score_input, "coefficients", "bias"],
            ["scores"],
            Vec::new(),
        ),
        make_node(
            "ArgMax",
            ["scores"],
            ["class_index"],
            vec![int_attribute("axis", 1), int_attribute("keepdims", 0)],
        ),
        make_node(
            "Gather",
            ["labels", "class_index"],
            ["label"],
            vec![int_attribute("axis", 0)],
        ),
    ]);
    assemble_model(
        GraphProto {
            node: nodes,
            name: "linear_score_classifier".into(),
            initializer,
            doc_string: String::new(),
            input: vec![make_value_info(
                "X",
                &[
                    Dimension::Symbolic("batch".into()),
                    Dimension::Fixed(features),
                ],
            )],
            output: vec![make_typed_value_info(
                "label",
                &[Dimension::Symbolic("batch".into())],
                INT64,
            )],
            value_info: vec![make_value_info(
                "scores",
                &[
                    Dimension::Symbolic("batch".into()),
                    Dimension::Fixed(classes),
                ],
            )],
        },
        OPSET_VERSION,
        IR_VERSION,
    )
}