onnx-export-rs 0.1.1

Export canonical Rust machine-learning models to ONNX
Documentation
use crate::canonical::LogisticModelWeights;
use crate::graph_builder::{
    assemble_model, int_attribute, make_node, make_tensor, make_value_info, Dimension,
};
use crate::proto::{GraphProto, ModelProto};
use crate::{IR_VERSION, OPSET_VERSION};

/// Exports binary logistic regression as `Gemm` + `Sigmoid` and multiclass
/// logistic regression as `Gemm` + `Softmax(axis=1)`.
#[must_use]
pub fn export_logistic(weights: &LogisticModelWeights) -> ModelProto {
    let score_count = weights.coefficients.nrows();
    let features = weights.n_features();
    let coefficients = weights
        .coefficients
        .t()
        .mapv(|value| value as f32)
        .into_dyn();
    let intercept = weights.intercept.mapv(|value| value as f32).into_dyn();
    let (activation, attributes) = if weights.n_classes == 2 {
        ("Sigmoid", Vec::new())
    } else {
        ("Softmax", vec![int_attribute("axis", 1)])
    };
    let graph = GraphProto {
        node: vec![
            make_node("Gemm", ["X", "W", "b"], ["scores"], Vec::new()),
            make_node(activation, ["scores"], ["probabilities"], attributes),
        ],
        name: "logistic_regression".into(),
        initializer: vec![
            make_tensor("W", &coefficients),
            make_tensor("b", &intercept),
        ],
        doc_string: String::new(),
        input: vec![make_value_info(
            "X",
            &[
                Dimension::Symbolic("batch".into()),
                Dimension::Fixed(features),
            ],
        )],
        output: vec![make_value_info(
            "probabilities",
            &[
                Dimension::Symbolic("batch".into()),
                Dimension::Fixed(score_count),
            ],
        )],
        value_info: Vec::new(),
    };
    assemble_model(graph, OPSET_VERSION, IR_VERSION)
}