onnx-export-rs 0.1.1

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

use crate::canonical::LinearModelWeights;
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 linear regression as `Gemm(X, W, b)`.
#[must_use]
pub fn export_linear(weights: &LinearModelWeights) -> ModelProto {
    let features = weights.n_features();
    let coefficients = Array2::from_shape_fn((features, 1), |(feature, _)| {
        weights.coefficients[feature] as f32
    });
    let intercept = Array1::from_vec(vec![weights.intercept as f32]);
    let graph = GraphProto {
        node: vec![make_node("Gemm", ["X", "W", "b"], ["Y"], Vec::new())],
        name: "linear_regression".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)
}