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};
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)
}