onnx-export-rs 0.1.1

Export canonical Rust machine-learning models to ONNX
Documentation
use crate::canonical::CategoricalNaiveBayes;
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 categorical Naive Bayes using per-feature table lookups.
#[must_use]
pub fn export_categorical_naive_bayes(model: &CategoricalNaiveBayes) -> ModelProto {
    let features = model.feature_log_probabilities.len();
    let classes = model.class_labels.len();
    let mut nodes = vec![make_node(
        "Cast",
        ["X"],
        ["categories"],
        vec![int_attribute("to", INT64 as i64)],
    )];
    let mut initializer = vec![
        make_tensor(
            "log_priors",
            &model.log_priors.mapv(|value| value as f32).into_dyn(),
        ),
        make_i64_tensor("labels", &[classes], model.class_labels.clone()),
    ];
    let mut sum = "log_priors".to_owned();
    for (feature, table) in model.feature_log_probabilities.iter().enumerate() {
        let index_name = format!("feature_index_{feature}");
        let table_name = format!("table_{feature}");
        let category_name = format!("category_{feature}");
        let score_name = format!("feature_score_{feature}");
        let sum_name = format!("score_sum_{feature}");
        initializer.push(make_i64_tensor(&index_name, &[], vec![feature as i64]));
        initializer.push(make_tensor(
            &table_name,
            &table.mapv(|value| value as f32).into_dyn(),
        ));
        nodes.extend([
            make_node(
                "Gather",
                ["categories", &index_name],
                [&category_name],
                vec![int_attribute("axis", 1)],
            ),
            make_node(
                "Gather",
                [&table_name, &category_name],
                [&score_name],
                vec![int_attribute("axis", 0)],
            ),
            make_node("Add", [&sum, &score_name], [&sum_name], Vec::new()),
        ]);
        sum = sum_name;
    }
    nodes.extend([
        make_node(
            "ArgMax",
            [&sum],
            ["class_index"],
            vec![int_attribute("axis", 1), int_attribute("keepdims", 0)],
        ),
        make_node(
            "Gather",
            ["labels", "class_index"],
            ["label"],
            vec![int_attribute("axis", 0)],
        ),
    ]);
    assemble_model(
        GraphProto {
            node: nodes,
            name: "categorical_naive_bayes".into(),
            initializer,
            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(),
        },
        OPSET_VERSION,
        IR_VERSION,
    )
}