use crate::data::DMatrix;
use crate::error::HessboostError;
use crate::tree::linear::{LinearLeaves, UncheckedLinearLeaves};
use crate::tree::{SplitTest, split_goes_left};
use serde::{Deserialize, Serialize};
const NO_CHILD: i32 = -1;
#[derive(Debug, Clone, Copy, PartialEq, Serialize, Deserialize)]
#[non_exhaustive]
pub struct Node {
pub split_feature: u32,
pub split_cond: f32,
pub default_left: bool,
pub left: i32,
pub right: i32,
pub leaf_value: f32,
pub sum_hess: f32,
pub split_gain: f32,
pub is_categorical: bool,
pub cat_begin: u32,
pub cat_end: u32,
}
impl Node {
pub(crate) fn leaf(value: f32, sum_hess: f32) -> Self {
Node {
split_feature: 0,
split_cond: 0.0,
default_left: true,
left: NO_CHILD,
right: NO_CHILD,
leaf_value: value,
sum_hess,
split_gain: 0.0,
is_categorical: false,
cat_begin: 0,
cat_end: 0,
}
}
#[inline]
pub fn is_leaf(&self) -> bool {
self.left == NO_CHILD
}
}
#[derive(Debug, Clone, Copy)]
pub(crate) struct SplitRule<'a> {
feature: u32,
test: SplitTest<'a>,
default_left: bool,
}
impl<'a> SplitRule<'a> {
pub(crate) fn numeric(feature: u32, threshold: f32, default_left: bool) -> Self {
SplitRule {
feature,
test: SplitTest::Threshold(threshold),
default_left,
}
}
pub(crate) fn categorical(feature: u32, cats_left: &'a [u32], default_left: bool) -> Self {
SplitRule {
feature,
test: SplitTest::Categories(cats_left),
default_left,
}
}
}
#[derive(Debug, Clone, Copy)]
pub(crate) struct ChildLeaf {
value: f32,
sum_hess: f32,
}
impl ChildLeaf {
pub(crate) fn new(value: f32, sum_hess: f32) -> Self {
ChildLeaf { value, sum_hess }
}
}
#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
#[serde(try_from = "UncheckedRegTree")]
pub struct RegTree {
nodes: Vec<Node>,
categories: Vec<u32>,
size_leaf_vector: usize,
leaf_vectors: Vec<f32>,
linear: Option<LinearLeaves>,
}
#[derive(Deserialize)]
pub(crate) struct UncheckedRegTree {
nodes: Vec<Node>,
categories: Vec<u32>,
size_leaf_vector: Option<usize>,
#[serde(default)]
leaf_vectors: Vec<f32>,
#[serde(deserialize_with = "Option::deserialize")]
linear: Option<UncheckedLinearLeaves>,
}
impl UncheckedRegTree {
pub(crate) fn states_leaf_width(&self) -> bool {
self.size_leaf_vector.is_some()
}
pub(crate) fn into_unchecked(self) -> RegTree {
RegTree::from_parts(
self.nodes,
self.categories,
self.size_leaf_vector.unwrap_or(0),
self.leaf_vectors,
self.linear.map(UncheckedLinearLeaves::into_unchecked),
)
}
}
impl TryFrom<UncheckedRegTree> for RegTree {
type Error = HessboostError;
fn try_from(unchecked: UncheckedRegTree) -> Result<Self, Self::Error> {
let tree = unchecked.into_unchecked();
if tree.is_valid_for_features(usize::MAX) {
Ok(tree)
} else {
Err(HessboostError::model_format("tree contains invalid nodes"))
}
}
}
impl RegTree {
pub(crate) fn with_root(sum_hess: f32) -> Self {
Self::from_scalar_parts(vec![Node::leaf(0.0, sum_hess)], Vec::new())
}
pub(crate) fn from_scalar_parts(nodes: Vec<Node>, categories: Vec<u32>) -> Self {
RegTree {
nodes,
categories,
size_leaf_vector: 0,
leaf_vectors: Vec::new(),
linear: None,
}
}
pub(crate) fn from_parts(
nodes: Vec<Node>,
categories: Vec<u32>,
size_leaf_vector: usize,
leaf_vectors: Vec<f32>,
linear: Option<LinearLeaves>,
) -> Self {
RegTree {
nodes,
categories,
size_leaf_vector,
leaf_vectors,
linear,
}
}
pub(crate) fn leaf_vector_parts(&self) -> (usize, &[f32]) {
(self.size_leaf_vector, &self.leaf_vectors)
}
pub(crate) fn with_vector_root(n_outputs: usize, sum_hess: f32) -> Self {
debug_assert!(n_outputs > 1);
RegTree {
nodes: vec![Node::leaf(0.0, sum_hess)],
categories: Vec::new(),
size_leaf_vector: n_outputs,
leaf_vectors: vec![0.0; n_outputs],
linear: None,
}
}
#[inline]
pub fn size_leaf_vector(&self) -> usize {
self.size_leaf_vector.max(1)
}
#[inline]
pub fn is_vector_leaf(&self) -> bool {
self.size_leaf_vector > 1
}
#[inline]
pub fn leaf_vector(&self, nid: usize) -> &[f32] {
if self.is_vector_leaf() {
let k = self.size_leaf_vector;
&self.leaf_vectors[nid * k..(nid + 1) * k]
} else {
std::slice::from_ref(&self.nodes[nid].leaf_value)
}
}
pub(crate) fn set_leaf_vector(&mut self, nid: usize, values: &[f32]) {
let k = self.size_leaf_vector;
debug_assert!(k > 1 && values.len() == k);
self.leaf_vectors[nid * k..(nid + 1) * k].copy_from_slice(values);
}
fn grow_leaf_vectors(&mut self) {
if self.is_vector_leaf() {
self.leaf_vectors
.resize(self.nodes.len() * self.size_leaf_vector, 0.0);
}
}
#[inline]
pub fn num_nodes(&self) -> usize {
self.nodes.len()
}
pub fn num_leaves(&self) -> usize {
self.nodes.iter().filter(|n| n.is_leaf()).count()
}
pub(crate) fn is_valid_for_features(&self, n_features: usize) -> bool {
let locally_valid = !self.nodes.is_empty()
&& self.nodes.iter().all(|node| {
node.sum_hess.is_finite()
&& node.leaf_value.is_finite()
&& node.split_cond.is_finite()
&& node.split_gain.is_finite()
&& ((node.is_leaf() && node.right == NO_CHILD)
|| ((node.split_feature as usize) < n_features
&& node.left >= 0
&& node.right >= 0
&& (node.left as usize) < self.nodes.len()
&& (node.right as usize) < self.nodes.len()))
&& (!node.is_categorical
|| (node.is_leaf() && self.is_vector_leaf())
|| (node.cat_begin < node.cat_end
&& (node.cat_end as usize) <= self.categories.len()))
})
&& self
.linear
.as_ref()
.is_none_or(|linear| linear.is_valid(&self.nodes, n_features));
if !locally_valid {
return false;
}
if self.is_vector_leaf()
&& (self.nodes.len().checked_mul(self.size_leaf_vector)
!= Some(self.leaf_vectors.len())
|| self.leaf_vectors.iter().any(|w| !w.is_finite()))
{
return false;
}
if self.size_leaf_vector == 1
|| (!self.is_vector_leaf() && !self.leaf_vectors.is_empty())
|| (self.is_vector_leaf() && self.linear.is_some())
{
return false;
}
let mut seen = vec![false; self.nodes.len()];
let mut stack = vec![0usize];
while let Some(node_id) = stack.pop() {
if seen[node_id] {
return false;
}
seen[node_id] = true;
let node = &self.nodes[node_id];
if !node.is_leaf() {
stack.push(node.left as usize);
stack.push(node.right as usize);
}
}
seen.into_iter().all(|visited| visited)
}
#[inline]
pub fn nodes(&self) -> &[Node] {
&self.nodes
}
#[inline]
pub(crate) fn categories(&self) -> &[u32] {
&self.categories
}
#[inline]
pub(crate) fn node_categories(&self, node: &Node) -> &[u32] {
&self.categories[node.cat_begin as usize..node.cat_end as usize]
}
#[inline]
pub fn linear_leaves(&self) -> Option<&LinearLeaves> {
self.linear.as_ref()
}
pub(crate) fn set_linear_leaves(&mut self, linear: LinearLeaves) {
self.linear = Some(linear);
}
#[inline]
pub fn node(&self, id: usize) -> &Node {
&self.nodes[id]
}
pub(crate) fn expand(
&mut self,
nid: usize,
split: SplitRule<'_>,
left: ChildLeaf,
right: ChildLeaf,
) -> (usize, usize) {
match split.test {
SplitTest::Threshold(cond) => self.nodes[nid].split_cond = cond,
SplitTest::Categories(cats_left) => {
let begin = self.categories.len() as u32;
self.categories.extend_from_slice(cats_left);
let n = &mut self.nodes[nid];
n.is_categorical = true;
n.cat_begin = begin;
n.cat_end = self.categories.len() as u32;
}
}
let left_id = self.nodes.len();
let right_id = left_id + 1;
let n = &mut self.nodes[nid];
n.split_feature = split.feature;
n.default_left = split.default_left;
n.left = left_id as i32;
n.right = right_id as i32;
self.nodes.push(Node::leaf(left.value, left.sum_hess));
self.nodes.push(Node::leaf(right.value, right.sum_hess));
self.grow_leaf_vectors();
(left_id, right_id)
}
pub(crate) fn set_leaf_value(&mut self, nid: usize, value: f32) {
self.nodes[nid].leaf_value = value;
}
pub(crate) fn set_split_gain(&mut self, nid: usize, gain: f32) {
self.nodes[nid].split_gain = gain;
}
pub(crate) fn set_sum_hess(&mut self, nid: usize, sum_hess: f32) {
self.nodes[nid].sum_hess = sum_hess;
}
pub(crate) fn scale_leaves(&mut self, factor: f32) {
let k = self.size_leaf_vector;
for (id, n) in self.nodes.iter_mut().enumerate() {
if n.is_leaf() {
n.leaf_value *= factor;
if k > 1 {
for w in &mut self.leaf_vectors[id * k..(id + 1) * k] {
*w *= factor;
}
}
}
}
if let Some(linear) = &mut self.linear {
linear.scale(f64::from(factor));
}
}
pub(crate) fn shift_leaves(&mut self, delta: f32) {
for n in &mut self.nodes {
if n.is_leaf() {
n.leaf_value += delta;
}
}
}
pub fn leaf_id_with(&self, get: impl Fn(u32) -> Option<f32>) -> usize {
let nodes = &self.nodes[..];
let mut nid = 0usize;
loop {
let node = &nodes[nid];
if node.is_leaf() {
return nid;
}
nid = if self.goes_left(node, get(node.split_feature)) {
node.left as usize
} else {
node.right as usize
};
}
}
#[inline]
pub(crate) fn child(&self, nid: usize, value: Option<f32>) -> usize {
let node = &self.nodes[nid];
if self.goes_left(node, value) {
node.left as usize
} else {
node.right as usize
}
}
#[inline]
pub(crate) fn goes_left(&self, node: &Node, value: Option<f32>) -> bool {
let test = if node.is_categorical {
SplitTest::Categories(self.node_categories(node))
} else {
SplitTest::Threshold(node.split_cond)
};
split_goes_left(value, node.default_left, test)
}
#[inline]
pub fn leaf_id_dense(&self, row: &[f32], missing: f32) -> usize {
self.leaf_id_with(|f| {
let v = row[f as usize];
(!crate::data::is_missing(v, missing)).then_some(v)
})
}
pub(crate) fn predict_row(&self, data: &DMatrix, row: usize) -> f32 {
debug_assert!(!self.is_vector_leaf(), "predict_row on a vector-leaf tree");
let get = |f: u32| data.get(row, f as usize);
let leaf = self.leaf_id_with(get);
let constant = self.nodes[leaf].leaf_value;
match &self.linear {
Some(linear) => linear.predict(leaf, constant, get),
None => constant,
}
}
}
#[cfg(test)]
mod tests {
use super::*;
fn stump() -> RegTree {
let mut t = RegTree::with_root(10.0);
t.expand(
0,
SplitRule::numeric(0, 0.5, true),
ChildLeaf::new(-1.0, 5.0),
ChildLeaf::new(2.0, 5.0),
);
t
}
#[test]
fn deserialization_refuses_malformed_trees() {
let doc = serde_json::to_value(stump()).unwrap();
assert_eq!(
serde_json::from_value::<RegTree>(doc.clone()).unwrap(),
stump()
);
for (node, field, value) in [(0, "left", 3), (0, "right", 0), (0, "left", 2)] {
let mut bad = doc.clone();
bad["nodes"][node][field] = value.into();
assert!(
serde_json::from_value::<RegTree>(bad).is_err(),
"{field} = {value}"
);
}
let mut vector = serde_json::to_value(RegTree::with_vector_root(2, 1.0)).unwrap();
assert!(serde_json::from_value::<RegTree>(vector.clone()).is_ok());
vector["leaf_vectors"] = serde_json::json!([0.0]);
assert!(serde_json::from_value::<RegTree>(vector).is_err());
let mut wide = doc;
wide["size_leaf_vector"] = usize::MAX.into();
wide["leaf_vectors"] = serde_json::json!([]);
assert!(serde_json::from_value::<RegTree>(wide).is_err());
}
#[test]
fn routing_numeric() {
let t = stump();
assert_eq!(t.leaf_id_with(|_| Some(0.2)), 1);
assert_eq!(t.leaf_id_with(|_| Some(0.9)), 2);
}
#[test]
fn routing_missing_follows_default() {
let t = stump();
assert_eq!(t.leaf_id_with(|_| None), 1);
assert_eq!(t.node(1).leaf_value, -1.0);
}
#[test]
fn routing_categorical_set_membership() {
let mut t = RegTree::with_root(10.0);
t.expand(
0,
SplitRule::categorical(0, &[0, 2], false),
ChildLeaf::new(-1.0, 5.0),
ChildLeaf::new(2.0, 5.0),
);
assert!(t.node(0).is_categorical);
assert_eq!(t.leaf_id_with(|_| Some(0.0)), 1);
assert_eq!(t.leaf_id_with(|_| Some(2.0)), 1);
assert_eq!(t.leaf_id_with(|_| Some(1.0)), 2);
assert_eq!(t.leaf_id_with(|_| Some(3.0)), 2);
assert_eq!(t.leaf_id_with(|_| Some(9.0)), 2);
assert_eq!(t.leaf_id_with(|_| None), 2);
}
#[test]
fn predict_row_dense() {
let t = stump();
let d = DMatrix::from_dense(&[0.1, 0.9], 2, 1).unwrap();
assert_eq!(t.predict_row(&d, 0), -1.0);
assert_eq!(t.predict_row(&d, 1), 2.0);
assert_eq!(t.num_leaves(), 2);
assert_eq!(t.num_nodes(), 3);
}
#[test]
fn vector_leaves_refuse_linear_payload() {
let linear: LinearLeaves = serde_json::from_str(
r#"{"offsets":[0,1],"intercepts":[0.5],"features":[0],"coeffs":[2.0]}"#,
)
.unwrap();
let mut scalar = RegTree::with_root(1.0);
scalar.set_linear_leaves(linear.clone());
assert!(scalar.is_valid_for_features(1));
let mut vector = RegTree::with_vector_root(2, 1.0);
assert!(vector.is_valid_for_features(1));
vector.set_linear_leaves(linear);
assert!(!vector.is_valid_for_features(1));
}
#[test]
fn a_leaf_refuses_a_right_child() {
let mut tree = RegTree::with_root(1.0);
assert!(tree.is_valid_for_features(1));
tree.nodes[0].right = -2;
assert!(!tree.is_valid_for_features(1));
}
}