onnx-export-rs 0.1.1

Export canonical Rust machine-learning models to ONNX
Documentation
use crate::canonical::{ForestStructure, TreeTask};
use crate::graph_builder::{
    assemble_model, floats_attribute, int_attribute, ints_attribute, make_node,
    make_typed_value_info, make_value_info, string_attribute, strings_attribute, Dimension, INT64,
};
use crate::proto::{ModelProto, OperatorSetIdProto};
use crate::{Error, Result, IR_VERSION, ML_OPSET_VERSION, OPSET_VERSION};

/// Exports a forest through ONNX-ML `TreeEnsembleRegressor` v3.
///
/// For [`TreeTask::Classification`], one-hot leaf scores are followed by
/// `ArgMax(axis=1)`, producing a zero-based `int64` class index.
pub fn export_tree_ensemble(forest: &ForestStructure, task: TreeTask) -> Result<ModelProto> {
    export_tree_ensemble_with_options(forest, task, &[], b"NONE")
}

pub(crate) fn export_tree_ensemble_with_options(
    forest: &ForestStructure,
    task: TreeTask,
    base_values: &[f32],
    post_transform: &[u8],
) -> Result<ModelProto> {
    if forest.trees.is_empty() || forest.n_targets == 0 {
        return Err(Error::InvalidModel(
            "a forest needs trees and targets".into(),
        ));
    }
    let mut tree_ids = Vec::new();
    let mut node_ids = Vec::new();
    let mut feature_ids = Vec::new();
    let mut values = Vec::new();
    let mut true_ids = Vec::new();
    let mut false_ids = Vec::new();
    let mut modes = Vec::new();
    let mut target_tree_ids = Vec::new();
    let mut target_node_ids = Vec::new();
    let mut target_ids = Vec::new();
    let mut target_weights = Vec::new();

    for (tree_index, tree) in forest.trees.iter().enumerate() {
        let tree_id =
            i64::try_from(tree_index).map_err(|_| Error::InvalidModel("too many trees".into()))?;
        if tree.nodes.is_empty() {
            return Err(Error::InvalidModel("tree has no nodes".into()));
        }
        for (expected, node) in tree.nodes.iter().enumerate() {
            if node.id != expected as i64 {
                return Err(Error::InvalidModel(
                    "tree node IDs must be sequential".into(),
                ));
            }
            tree_ids.push(tree_id);
            node_ids.push(node.id);
            feature_ids.push(node.feature_id);
            values.push(node.threshold);
            true_ids.push(node.true_child_id);
            false_ids.push(node.false_child_id);
            modes.push(if node.is_leaf() {
                b"LEAF".to_vec()
            } else {
                node.branch_mode.as_onnx().to_vec()
            });
            if node.is_leaf() {
                if node.leaf_values.len() != forest.n_targets {
                    return Err(Error::InvalidModel(format!(
                        "leaf {} has {} values; expected {}",
                        node.id,
                        node.leaf_values.len(),
                        forest.n_targets
                    )));
                }
                for (target, &weight) in node.leaf_values.iter().enumerate() {
                    target_tree_ids.push(tree_id);
                    target_node_ids.push(node.id);
                    target_ids.push(target as i64);
                    target_weights.push(weight);
                }
            }
        }
    }

    let mut attributes = vec![
        string_attribute(
            "aggregate_function",
            forest.aggregation.as_onnx().as_bytes(),
        ),
        int_attribute("n_targets", forest.n_targets as i64),
        ints_attribute("nodes_treeids", tree_ids),
        ints_attribute("nodes_nodeids", node_ids),
        ints_attribute("nodes_featureids", feature_ids),
        floats_attribute("nodes_values", values),
        ints_attribute("nodes_truenodeids", true_ids),
        ints_attribute("nodes_falsenodeids", false_ids),
        strings_attribute("nodes_modes", modes),
        ints_attribute("target_treeids", target_tree_ids),
        ints_attribute("target_nodeids", target_node_ids),
        ints_attribute("target_ids", target_ids),
        floats_attribute("target_weights", target_weights),
        string_attribute("post_transform", post_transform),
    ];
    if !base_values.is_empty() {
        attributes.push(floats_attribute("base_values", base_values.to_vec()));
    }
    let feature_count = forest
        .trees
        .iter()
        .flat_map(|tree| &tree.nodes)
        .filter(|node| !node.is_leaf())
        .map(|node| node.feature_id as usize + 1)
        .max()
        .unwrap_or(1);
    let tree_output = if task == TreeTask::Classification {
        "scores"
    } else {
        "Y"
    };
    let mut tree_node = make_node("TreeEnsembleRegressor", ["X"], [tree_output], attributes);
    tree_node.domain = "ai.onnx.ml".into();
    let mut nodes = vec![tree_node];
    let (outputs, value_info) = if task == TreeTask::Classification {
        nodes.push(make_node(
            "ArgMax",
            ["scores"],
            ["class_index"],
            vec![int_attribute("axis", 1), int_attribute("keepdims", 0)],
        ));
        (
            vec![make_typed_value_info(
                "class_index",
                &[Dimension::Symbolic("batch".into())],
                INT64,
            )],
            vec![make_value_info(
                "scores",
                &[
                    Dimension::Symbolic("batch".into()),
                    Dimension::Fixed(forest.n_targets),
                ],
            )],
        )
    } else {
        (
            vec![make_value_info(
                "Y",
                &[
                    Dimension::Symbolic("batch".into()),
                    Dimension::Fixed(forest.n_targets),
                ],
            )],
            Vec::new(),
        )
    };
    let graph = crate::proto::GraphProto {
        node: nodes,
        name: "tree_ensemble".into(),
        initializer: Vec::new(),
        doc_string: String::new(),
        input: vec![make_value_info(
            "X",
            &[
                Dimension::Symbolic("batch".into()),
                Dimension::Fixed(feature_count),
            ],
        )],
        output: outputs,
        value_info,
    };
    let mut model = assemble_model(graph, OPSET_VERSION, IR_VERSION);
    model.opset_import.push(OperatorSetIdProto {
        domain: "ai.onnx.ml".into(),
        version: ML_OPSET_VERSION,
    });
    Ok(model)
}