onnx-export-rs 0.1.1

Export canonical Rust machine-learning models to ONNX
Documentation
use ndarray::{Array1, Array2};

use crate::canonical::{GeneralizedLinearModel, LinkFunction};
use crate::graph_builder::{assemble_model, make_node, make_tensor, make_value_info, Dimension};
use crate::proto::{GraphProto, ModelProto};
use crate::{IR_VERSION, OPSET_VERSION};

/// Exports a generalized linear model as `Gemm(X, W, b)` followed by the
/// inverse-link activation: nothing for `Identity`, `Exp` for `Log`, and
/// `Sigmoid` for `Logit`.
#[must_use]
pub fn export_generalized_linear(model: &GeneralizedLinearModel) -> ModelProto {
    let features = model.n_features();
    let coefficients = Array2::from_shape_fn((features, 1), |(feature, _)| {
        model.coefficients[feature] as f32
    });
    let intercept = Array1::from_vec(vec![model.intercept as f32]);

    // The `Gemm` always writes `scores`; the activation (if any) renames it to
    // `Y`, so a bare identity link collapses to a single node emitting `Y`.
    let (gemm_output, mut nodes) = match model.link {
        LinkFunction::Identity => ("Y", Vec::new()),
        LinkFunction::Log => (
            "scores",
            vec![make_node("Exp", ["scores"], ["Y"], Vec::new())],
        ),
        LinkFunction::Logit => (
            "scores",
            vec![make_node("Sigmoid", ["scores"], ["Y"], Vec::new())],
        ),
    };
    nodes.insert(
        0,
        make_node("Gemm", ["X", "W", "b"], [gemm_output], Vec::new()),
    );

    let graph = GraphProto {
        node: nodes,
        name: "generalized_linear_model".into(),
        initializer: vec![
            make_tensor("W", &coefficients.into_dyn()),
            make_tensor("b", &intercept.into_dyn()),
        ],
        doc_string: String::new(),
        input: vec![make_value_info(
            "X",
            &[
                Dimension::Symbolic("batch".into()),
                Dimension::Fixed(features),
            ],
        )],
        output: vec![make_value_info(
            "Y",
            &[Dimension::Symbolic("batch".into()), Dimension::Fixed(1)],
        )],
        value_info: Vec::new(),
    };
    assemble_model(graph, OPSET_VERSION, IR_VERSION)
}