use crate::{Error, Result};
#[derive(Clone, Debug, PartialEq)]
pub enum RecursiveNode {
Branch {
feature: usize,
threshold: f64,
left: Box<Self>,
right: Box<Self>,
},
BranchLessThan {
feature: usize,
threshold: f64,
left: Box<Self>,
right: Box<Self>,
},
Leaf(Vec<f64>),
}
#[derive(Clone, Debug, PartialEq)]
pub struct TreeNode {
pub id: i64,
pub feature_id: i64,
pub threshold: f32,
pub true_child_id: i64,
pub false_child_id: i64,
pub branch_mode: BranchMode,
pub leaf_values: Vec<f32>,
}
#[derive(Clone, Copy, Debug, Default, PartialEq, Eq)]
pub enum BranchMode {
#[default]
LessOrEqual,
LessThan,
}
impl BranchMode {
pub(crate) const fn as_onnx(self) -> &'static [u8] {
match self {
Self::LessOrEqual => b"BRANCH_LEQ",
Self::LessThan => b"BRANCH_LT",
}
}
}
impl TreeNode {
#[must_use]
pub fn is_leaf(&self) -> bool {
!self.leaf_values.is_empty()
}
}
#[derive(Clone, Debug, PartialEq)]
pub struct TreeStructure {
pub nodes: Vec<TreeNode>,
}
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub enum AggregationMode {
Sum,
Average,
Min,
Max,
}
impl AggregationMode {
pub(crate) const fn as_onnx(self) -> &'static str {
match self {
Self::Sum => "SUM",
Self::Average => "AVERAGE",
Self::Min => "MIN",
Self::Max => "MAX",
}
}
}
#[derive(Clone, Debug, PartialEq)]
pub struct ForestStructure {
pub trees: Vec<TreeStructure>,
pub aggregation: AggregationMode,
pub n_targets: usize,
}
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub enum TreeTask {
Regression,
Classification,
}
#[derive(Clone, Copy, Debug, Default, PartialEq, Eq)]
pub enum PostTransform {
#[default]
None,
Logistic,
Softmax,
}
impl PostTransform {
pub(crate) const fn as_onnx(self) -> &'static [u8] {
match self {
Self::None => b"NONE",
Self::Logistic => b"LOGISTIC",
Self::Softmax => b"SOFTMAX",
}
}
}
#[derive(Clone, Debug, PartialEq)]
pub struct GradientBoostedEnsemble {
pub trees: Vec<TreeStructure>,
pub base_values: Vec<f64>,
pub learning_rate: f64,
pub n_targets: usize,
pub task: TreeTask,
pub post_transform: PostTransform,
}
pub fn flatten_tree(root: &RecursiveNode) -> Result<TreeStructure> {
fn visit(node: &RecursiveNode, output: &mut Vec<TreeNode>) -> Result<i64> {
let id = i64::try_from(output.len())
.map_err(|_| Error::InvalidModel("tree has too many nodes".into()))?;
output.push(TreeNode {
id,
feature_id: 0,
threshold: 0.0,
true_child_id: 0,
false_child_id: 0,
branch_mode: BranchMode::LessOrEqual,
leaf_values: Vec::new(),
});
match node {
RecursiveNode::Leaf(values) => {
if values.is_empty() {
return Err(Error::InvalidModel("tree leaf has no values".into()));
}
output[id as usize].leaf_values = values.iter().map(|&v| v as f32).collect();
}
RecursiveNode::Branch {
feature,
threshold,
left,
right,
}
| RecursiveNode::BranchLessThan {
feature,
threshold,
left,
right,
} => {
let true_id = visit(left, output)?;
let false_id = visit(right, output)?;
let flat = &mut output[id as usize];
flat.feature_id = i64::try_from(*feature)
.map_err(|_| Error::InvalidModel("feature index is too large".into()))?;
flat.threshold = *threshold as f32;
flat.true_child_id = true_id;
flat.false_child_id = false_id;
flat.branch_mode = if matches!(node, RecursiveNode::BranchLessThan { .. }) {
BranchMode::LessThan
} else {
BranchMode::LessOrEqual
};
}
}
Ok(id)
}
let mut nodes = Vec::new();
visit(root, &mut nodes)?;
Ok(TreeStructure { nodes })
}