onnx-export-rs 0.1.1

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

use crate::canonical::{KnnClassifier, KnnRegressor, KnnWeight};
use crate::graph_builder::{
    assemble_model, int_attribute, ints_attribute, make_i64_tensor, make_node, make_tensor,
    make_typed_value_info, make_value_info, Dimension, INT64,
};
use crate::proto::{GraphProto, ModelProto, NodeProto, TensorProto};
use crate::{IR_VERSION, OPSET_VERSION};

fn neighbor_graph(samples: &ndarray::Array2<f64>, k: usize) -> (Vec<NodeProto>, Vec<TensorProto>) {
    let sample_count = samples.nrows();
    let features = samples.ncols();
    let samples = Array3::from_shape_fn((1, sample_count, features), |(_, row, column)| {
        samples[(row, column)] as f32
    });
    (
        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("Sqrt", ["squared_distances"], ["distances"], Vec::new()),
            make_node("Neg", ["distances"], ["negative_distances"], Vec::new()),
            make_node(
                "TopK",
                ["negative_distances", "k"],
                ["negative_neighbor_distances", "indices"],
                vec![
                    int_attribute("axis", 1),
                    int_attribute("largest", 1),
                    int_attribute("sorted", 1),
                ],
            ),
            make_node(
                "Neg",
                ["negative_neighbor_distances"],
                ["neighbor_distances"],
                Vec::new(),
            ),
        ],
        vec![
            make_tensor("samples", &samples.into_dyn()),
            make_i64_tensor("axis_one", &[1], vec![1]),
            make_i64_tensor("axis_features", &[1], vec![2]),
            make_i64_tensor("k", &[1], vec![k as i64]),
        ],
    )
}

fn weight_nodes(weight: KnnWeight) -> (Vec<NodeProto>, Vec<TensorProto>, &'static str) {
    if weight == KnnWeight::Uniform {
        return (Vec::new(), Vec::new(), "");
    }
    let zero = Array1::from_vec(vec![0.0_f32]);
    let one = Array1::from_vec(vec![1.0_f32]);
    (
        vec![
            make_node(
                "Equal",
                ["neighbor_distances", "zero"],
                ["zero_mask"],
                Vec::new(),
            ),
            make_node(
                "Cast",
                ["zero_mask"],
                ["zero_weights"],
                vec![int_attribute("to", 1)],
            ),
            make_node(
                "ReduceSum",
                ["zero_weights", "axis_neighbors"],
                ["zero_count"],
                vec![int_attribute("keepdims", 1)],
            ),
            make_node("Greater", ["zero_count", "zero"], ["has_zero"], Vec::new()),
            make_node(
                "Where",
                ["zero_mask", "one", "neighbor_distances"],
                ["safe_distances"],
                Vec::new(),
            ),
            make_node(
                "Div",
                ["one", "safe_distances"],
                ["inverse_weights"],
                Vec::new(),
            ),
            make_node(
                "Where",
                ["has_zero", "zero_weights", "inverse_weights"],
                ["weights"],
                Vec::new(),
            ),
        ],
        vec![
            make_tensor("zero", &zero.into_dyn()),
            make_tensor("one", &one.into_dyn()),
            make_i64_tensor("axis_neighbors", &[1], vec![1]),
        ],
        "weights",
    )
}

/// Exports Euclidean k-nearest-neighbor regression.
#[must_use]
pub fn export_knn_regressor(model: &KnnRegressor) -> ModelProto {
    let (mut nodes, mut initializer) = neighbor_graph(&model.samples, model.k);
    let targets = model.targets.mapv(|value| value as f32);
    initializer.push(make_tensor("targets", &targets.into_dyn()));
    nodes.push(make_node(
        "Gather",
        ["targets", "indices"],
        ["neighbor_targets"],
        vec![int_attribute("axis", 0)],
    ));
    if model.weight == KnnWeight::Uniform {
        nodes.push(make_node(
            "ReduceMean",
            ["neighbor_targets"],
            ["prediction"],
            vec![
                ints_attribute("axes", vec![1]),
                int_attribute("keepdims", 1),
            ],
        ));
    } else {
        let (weighting, extra, _) = weight_nodes(model.weight);
        nodes.extend(weighting);
        initializer.extend(extra);
        nodes.extend([
            make_node(
                "Mul",
                ["neighbor_targets", "weights"],
                ["weighted_targets"],
                Vec::new(),
            ),
            make_node(
                "ReduceSum",
                ["weighted_targets", "axis_neighbors"],
                ["target_sum"],
                vec![int_attribute("keepdims", 1)],
            ),
            make_node(
                "ReduceSum",
                ["weights", "axis_neighbors"],
                ["weight_sum"],
                vec![int_attribute("keepdims", 1)],
            ),
            make_node(
                "Div",
                ["target_sum", "weight_sum"],
                ["prediction"],
                Vec::new(),
            ),
        ]);
    }
    assemble_knn(
        model.samples.ncols(),
        nodes,
        initializer,
        make_value_info(
            "prediction",
            &[Dimension::Symbolic("batch".into()), Dimension::Fixed(1)],
        ),
    )
}

/// Exports Euclidean k-nearest-neighbor classification.
#[must_use]
pub fn export_knn_classifier(model: &KnnClassifier) -> ModelProto {
    let (mut nodes, mut initializer) = neighbor_graph(&model.samples, model.k);
    let classes = model.class_labels.len();
    initializer.extend([
        make_i64_tensor(
            "targets",
            &[model.target_indices.len()],
            model.target_indices.clone(),
        ),
        make_i64_tensor("labels", &[classes], model.class_labels.clone()),
        make_i64_tensor("depth", &[], vec![classes as i64]),
        make_tensor(
            "one_hot_values",
            &Array1::from_vec(vec![0.0_f32, 1.0]).into_dyn(),
        ),
        make_i64_tensor("axis_neighbors", &[1], vec![1]),
    ]);
    nodes.extend([
        make_node(
            "Gather",
            ["targets", "indices"],
            ["neighbor_classes"],
            vec![int_attribute("axis", 0)],
        ),
        make_node(
            "OneHot",
            ["neighbor_classes", "depth", "one_hot_values"],
            ["votes"],
            vec![int_attribute("axis", -1)],
        ),
    ]);
    if model.weight == KnnWeight::Distance {
        let (weighting, extra, _) = weight_nodes(model.weight);
        nodes.extend(weighting);
        initializer.extend(extra);
        nodes.extend([
            make_node(
                "Unsqueeze",
                ["weights", "axis_two"],
                ["expanded_weights"],
                Vec::new(),
            ),
            make_node(
                "Mul",
                ["votes", "expanded_weights"],
                ["weighted_votes"],
                Vec::new(),
            ),
            make_node(
                "ReduceSum",
                ["weighted_votes", "axis_neighbors"],
                ["class_scores"],
                vec![int_attribute("keepdims", 0)],
            ),
        ]);
        initializer.push(make_i64_tensor("axis_two", &[1], vec![2]));
    } else {
        nodes.push(make_node(
            "ReduceSum",
            ["votes", "axis_neighbors"],
            ["class_scores"],
            vec![int_attribute("keepdims", 0)],
        ));
    }
    nodes.extend([
        make_node(
            "ArgMax",
            ["class_scores"],
            ["class_index"],
            vec![int_attribute("axis", 1), int_attribute("keepdims", 0)],
        ),
        make_node(
            "Gather",
            ["labels", "class_index"],
            ["label"],
            vec![int_attribute("axis", 0)],
        ),
    ]);
    assemble_knn(
        model.samples.ncols(),
        nodes,
        initializer,
        make_typed_value_info("label", &[Dimension::Symbolic("batch".into())], INT64),
    )
}

fn assemble_knn(
    features: usize,
    nodes: Vec<NodeProto>,
    initializer: Vec<TensorProto>,
    output: crate::proto::ValueInfoProto,
) -> ModelProto {
    assemble_model(
        GraphProto {
            node: nodes,
            name: "k_nearest_neighbors".into(),
            initializer,
            doc_string: String::new(),
            input: vec![make_value_info(
                "X",
                &[
                    Dimension::Symbolic("batch".into()),
                    Dimension::Fixed(features),
                ],
            )],
            output: vec![output],
            value_info: Vec::new(),
        },
        OPSET_VERSION,
        IR_VERSION,
    )
}