onnx-export-rs 0.1.1

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

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

/// Exports full-covariance Gaussian-mixture component assignment.
#[must_use]
pub fn export_gaussian_mixture(model: &GaussianMixture) -> ModelProto {
    let components = model.means.nrows();
    let features = model.means.ncols();
    let means = Array3::from_shape_fn((1, components, features), |(_, component, feature)| {
        model.means[(component, feature)] as f32
    });
    let precisions = model.precisions.mapv(|value| value as f32);
    let offsets = model.offsets.mapv(|value| value as f32);
    let minus_half = Array1::from_vec(vec![-0.5_f32]);
    let graph = GraphProto {
        node: vec![
            make_node("Unsqueeze", ["X", "axis_one"], ["expanded"], Vec::new()),
            make_node("Sub", ["expanded", "means"], ["delta"], Vec::new()),
            make_node(
                "Einsum",
                ["delta", "precisions", "delta"],
                ["mahalanobis"],
                vec![string_attribute("equation", b"bcf,cfg,bcg->bc".to_vec())],
            ),
            make_node(
                "Mul",
                ["mahalanobis", "minus_half"],
                ["quadratic"],
                Vec::new(),
            ),
            make_node("Add", ["quadratic", "offsets"], ["scores"], Vec::new()),
            make_node(
                "ArgMax",
                ["scores"],
                ["component"],
                vec![int_attribute("axis", 1), int_attribute("keepdims", 0)],
            ),
        ],
        name: "gaussian_mixture".into(),
        initializer: vec![
            make_tensor("means", &means.into_dyn()),
            make_tensor("precisions", &precisions.into_dyn()),
            make_tensor("offsets", &offsets.into_dyn()),
            make_tensor("minus_half", &minus_half.into_dyn()),
            make_i64_tensor("axis_one", &[1], vec![1]),
        ],
        doc_string: String::new(),
        input: vec![make_value_info(
            "X",
            &[
                Dimension::Symbolic("batch".into()),
                Dimension::Fixed(features),
            ],
        )],
        output: vec![make_typed_value_info(
            "component",
            &[Dimension::Symbolic("batch".into())],
            INT64,
        )],
        value_info: Vec::new(),
    };
    assemble_model(graph, OPSET_VERSION, IR_VERSION)
}