onnx-export-rs 0.1.1

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

use crate::canonical::DbscanModel;
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 SmartCore-compatible DBSCAN radius-voting prediction.
#[must_use]
pub fn export_dbscan(model: &DbscanModel) -> ModelProto {
    let samples = model.samples.nrows();
    let features = model.samples.ncols();
    let classes_with_noise = model.n_clusters + 1;
    let sample_tensor = Array3::from_shape_fn((1, samples, features), |(_, row, column)| {
        model.samples[(row, column)] as f32
    });
    let votes = Array2::from_shape_fn((samples, classes_with_noise), |(sample, class)| {
        let fitted = model.cluster_indices[sample];
        let fitted_class = if fitted < 0 {
            model.n_clusters
        } else {
            fitted as usize
        };
        if class == fitted_class {
            1.0
        } else {
            0.0
        }
    });
    let graph = GraphProto {
        node: vec![
            make_node("Unsqueeze", ["X", "axis_one"], ["expanded"], Vec::new()),
            make_node("Sub", ["expanded", "samples"], ["delta"], Vec::new()),
            make_node("Mul", ["delta", "delta"], ["squared"], Vec::new()),
            make_node(
                "ReduceSum",
                ["squared", "axis_features"],
                ["squared_distances"],
                vec![int_attribute("keepdims", 0)],
            ),
            make_node(
                "LessOrEqual",
                ["squared_distances", "epsilon_squared"],
                ["neighbor_mask"],
                Vec::new(),
            ),
            make_node(
                "Cast",
                ["neighbor_mask"],
                ["neighbors"],
                vec![int_attribute("to", 1)],
            ),
            make_node("MatMul", ["neighbors", "votes"], ["counts"], Vec::new()),
            make_node(
                "ArgMax",
                ["counts"],
                ["class_index"],
                vec![int_attribute("axis", 1), int_attribute("keepdims", 0)],
            ),
            make_node(
                "Equal",
                ["class_index", "noise_index"],
                ["is_noise"],
                Vec::new(),
            ),
            make_node("Add", ["class_index", "one"], ["cluster_label"], Vec::new()),
            make_node(
                "Where",
                ["is_noise", "zero", "cluster_label"],
                ["label"],
                Vec::new(),
            ),
        ],
        name: "dbscan".into(),
        initializer: vec![
            make_tensor("samples", &sample_tensor.into_dyn()),
            make_tensor("votes", &votes.into_dyn()),
            make_tensor(
                "epsilon_squared",
                &Array1::from_vec(vec![(model.epsilon * model.epsilon) as f32]).into_dyn(),
            ),
            make_i64_tensor("axis_one", &[1], vec![1]),
            make_i64_tensor("axis_features", &[1], vec![2]),
            make_i64_tensor("noise_index", &[1], vec![model.n_clusters as i64]),
            make_i64_tensor("one", &[1], vec![1]),
            make_i64_tensor("zero", &[1], vec![0]),
        ],
        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)
}