onnx-export-rs 0.1.1

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

use crate::canonical::CentroidModel;
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 nearest-centroid prediction using squared Euclidean distance.
///
/// The `int64` output is the zero-based centroid index, matching SmartCore's
/// k-means prediction convention.
#[must_use]
pub fn export_nearest_centroid(model: &CentroidModel) -> ModelProto {
    let count = model.centroids.nrows();
    let features = model.n_features();
    let centroids = Array3::from_shape_fn((1, count, features), |(_, row, column)| {
        model.centroids[(row, column)] as f32
    });
    let graph = GraphProto {
        node: vec![
            make_node("Unsqueeze", ["X", "axis_one"], ["expanded"], Vec::new()),
            make_node("Sub", ["expanded", "centroids"], ["delta"], Vec::new()),
            make_node("Mul", ["delta", "delta"], ["squared"], Vec::new()),
            make_node(
                "ReduceSum",
                ["squared", "axis_features"],
                ["distances"],
                vec![int_attribute("keepdims", 0)],
            ),
            make_node(
                "ArgMin",
                ["distances"],
                ["label"],
                vec![int_attribute("axis", 1), int_attribute("keepdims", 0)],
            ),
        ],
        name: "nearest_centroid".into(),
        initializer: vec![
            make_tensor("centroids", &centroids.into_dyn()),
            make_i64_tensor("axis_one", &[1], vec![1]),
            make_i64_tensor("axis_features", &[1], vec![2]),
        ],
        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::new(),
    };
    assemble_model(graph, OPSET_VERSION, IR_VERSION)
}