pub(crate) mod grid;
use std::ops::Range;
use serde::{Deserialize, Serialize};
use crate::data::DMatrix;
use crate::error::{HessboostError, Result};
use crate::model::BoostedModel;
use crate::objective::Objective;
use crate::tree::RegTree;
use grid::TermGrid;
#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
#[non_exhaustive]
pub struct EbmInfo {
pub terms: Vec<Vec<u32>>,
pub tree_terms: Vec<u32>,
pub term_means: Vec<f64>,
#[serde(deserialize_with = "Option::deserialize")]
pub boulevard: Option<EbmBoulevard>,
}
#[derive(Debug, Clone, Copy, PartialEq, Serialize, Deserialize)]
#[non_exhaustive]
pub struct EbmBoulevard {
pub learning_rate: f64,
pub subsample: f64,
pub reg_lambda: f64,
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub(crate) struct Stage {
pub(crate) index: usize,
pub(crate) terms: Range<usize>,
pub(crate) trees: Range<usize>,
pub(crate) rounds: usize,
}
fn invalid_record(reason: impl std::fmt::Display) -> HessboostError {
HessboostError::model_format(format!("invalid EBM record: {reason}"))
}
impl EbmInfo {
pub(crate) fn validate(&self, model: &BoostedModel) -> Result<()> {
let fail = |reason: String| Err(invalid_record(reason));
if model.n_outputs() != 1
|| model.has_vector_leaves()
|| model.has_non_unit_tree_weights()
|| model.num_parallel_tree() != 1
|| model.trees().iter().any(|t| t.linear_leaves().is_some())
{
return fail("only single-output scalar tree ensembles are EBMs".into());
}
if self.tree_terms.len() != model.num_trees() {
return fail(format!(
"{} tree terms for {} trees",
self.tree_terms.len(),
model.num_trees()
));
}
if self.term_means.len() != self.terms.len() {
return fail(format!(
"{} term means for {} terms",
self.term_means.len(),
self.terms.len()
));
}
if self.term_means.iter().any(|m| !m.is_finite()) {
return fail("term means must be finite".into());
}
for (t, features) in self.terms.iter().enumerate() {
let sorted = features.windows(2).all(|w| w[0] < w[1]);
if !(1..=2).contains(&features.len())
|| !sorted
|| features.iter().any(|&f| f as usize >= model.n_features())
{
return fail(format!(
"term {t} must name one or two ascending features below {}",
model.n_features()
));
}
}
let mut categorical: Vec<[Option<bool>; 2]> = vec![[None; 2]; self.terms.len()];
for (i, (&term, tree)) in self.tree_terms.iter().zip(model.trees()).enumerate() {
let Some(features) = self.terms.get(term as usize) else {
return fail(format!(
"tree {i} names term {term} of {}",
self.terms.len()
));
};
for n in tree.nodes().iter().filter(|n| !n.is_leaf()) {
let Some(a) = features.iter().position(|&f| f == n.split_feature) else {
return fail(format!("tree {i} splits outside its term's features"));
};
let kind = &mut categorical[term as usize][a];
if *kind.get_or_insert(n.is_categorical) != n.is_categorical {
return fail(format!(
"term {term} splits feature {} both numerically and categorically",
n.split_feature
));
}
}
}
if let Some(b) = &self.boulevard {
if !(b.learning_rate > 0.0 && b.learning_rate <= 1.0) {
return fail("learning_rate must be in (0, 1]".into());
}
if !(b.subsample > 0.0 && b.subsample <= 1.0) {
return fail("subsample must be in (0, 1]".into());
}
if !(b.reg_lambda.is_finite() && b.reg_lambda >= 0.0) {
return fail("reg_lambda must be finite and >= 0".into());
}
if !model
.objective()
.built_in()
.is_some_and(Objective::is_unweighted_squared_error)
{
return fail("a Boulevard EBM is a reg:squarederror model".into());
}
self.stages().map(drop)?;
}
Ok(())
}
pub(crate) fn stages(&self) -> Result<impl Iterator<Item = Stage>> {
let mains = self.terms.iter().take_while(|t| t.len() == 1).count();
if self.terms[mains..].iter().any(|t| t.len() == 1) {
return Err(invalid_record(
"a Boulevard EBM lists its main terms before its pairs",
));
}
let mut at = 0;
let mut stages = Vec::with_capacity(2);
for (index, terms) in [0..mains, mains..self.terms.len()].into_iter().enumerate() {
let k = terms.len();
let trees = self.tree_terms[at..]
.iter()
.take_while(|&&t| terms.contains(&(t as usize)))
.count();
if k == 0 {
continue;
}
let round_robin = self.tree_terms[at..at + trees]
.iter()
.enumerate()
.all(|(i, &t)| t as usize == terms.start + i % k);
if trees % k != 0 || !round_robin {
return Err(invalid_record(format!(
"a Boulevard EBM stage must hold whole rounds of its {k} terms in term order"
)));
}
stages.push(Stage {
index,
terms,
trees: at..at + trees,
rounds: trees / k,
});
at += trees;
}
if at != self.tree_terms.len() {
return Err(invalid_record(
"a Boulevard EBM's trees must be its main stage, then its pair stage",
));
}
Ok(stages.into_iter())
}
pub(crate) fn term_trees<'m>(&self, model: &'m BoostedModel, term: usize) -> Vec<&'m RegTree> {
self.tree_terms
.iter()
.zip(model.trees())
.filter(|&(&t, _)| t as usize == term)
.map(|(_, tree)| tree)
.collect()
}
pub(crate) fn grid(&self, model: &BoostedModel, term: usize) -> TermGrid {
TermGrid::new(&self.term_trees(model, term), &self.terms[term])
}
pub(crate) fn term_means_on(&self, model: &BoostedModel, data: &DMatrix) -> Vec<f64> {
let n = data.n_rows();
(0..self.terms.len())
.map(|t| {
let grid = self.grid(model, t);
let values = self.raw_values(model, t, &grid);
let sum: f64 = (0..n).map(|row| values[grid.cell_of_row(data, row)]).sum();
sum / n as f64
})
.collect()
}
pub(crate) fn raw_values(
&self,
model: &BoostedModel,
term: usize,
grid: &TermGrid,
) -> Vec<f64> {
let trees = self.term_trees(model, term);
let mut leaf_values = Vec::with_capacity(grid.leaves.len());
for (t, tree) in trees.iter().enumerate() {
for leaf in &grid.leaves[grid.leaf_start[t]..grid.leaf_start[t + 1]] {
leaf_values.push(f64::from(tree.node(leaf.node as usize).leaf_value));
}
}
grid.paint(|i| leaf_values[i])
}
fn shape(&self, model: &BoostedModel, term: usize) -> TermShape {
let grid = self.grid(model, term);
let mean = self.term_means[term];
let values = self
.raw_values(model, term, &grid)
.into_iter()
.map(|v| v - mean)
.collect();
TermShape {
features: self.terms[term].clone(),
axes: grid.axes.into_iter().map(|a| a.kind).collect(),
values,
}
}
}
#[derive(Debug, Clone, PartialEq)]
#[non_exhaustive]
pub struct ShapeFunctions {
pub intercept: f64,
pub terms: Vec<TermShape>,
}
#[derive(Debug, Clone, PartialEq)]
#[non_exhaustive]
pub enum TermAxis {
#[non_exhaustive]
Numeric {
edges: Vec<f32>,
},
#[non_exhaustive]
Categorical {
categories: Vec<u32>,
},
}
impl TermAxis {
pub fn cells(&self) -> usize {
match self {
TermAxis::Numeric { edges } => edges.len() + 2,
TermAxis::Categorical { categories } => categories.len() + 2,
}
}
pub(crate) fn cell(&self, value: Option<f32>) -> usize {
match value {
Some(x) if !x.is_nan() => match self {
TermAxis::Numeric { edges } => edges.partition_point(|&e| e <= x),
TermAxis::Categorical { categories } => {
let code = x as u32;
categories.binary_search(&code).unwrap_or(categories.len())
}
},
_ => self.cells() - 1,
}
}
}
#[derive(Debug, Clone, PartialEq)]
pub struct TermShape {
features: Vec<u32>,
axes: Vec<TermAxis>,
values: Vec<f64>,
}
impl TermShape {
pub fn features(&self) -> &[u32] {
&self.features
}
pub fn axes(&self) -> &[TermAxis] {
&self.axes
}
pub fn values(&self) -> &[f64] {
&self.values
}
pub fn cell(&self, x: &[f32]) -> Result<usize> {
if x.len() != self.features.len() {
return Err(HessboostError::dimension_mismatch(
"term feature values",
self.features.len(),
x.len(),
));
}
Ok(self.axes.iter().zip(x).fold(0, |cell, (axis, &v)| {
cell * axis.cells() + axis.cell(Some(v))
}))
}
pub fn value(&self, x: &[f32]) -> Result<f64> {
Ok(self.values[self.cell(x)?])
}
}
fn ebm_info(model: &BoostedModel) -> Result<&EbmInfo> {
model.ebm().ok_or_else(|| {
HessboostError::incompatible_model("model", "not an EBM: train it with `booster = ebm`")
})
}
pub fn shape_functions(model: &BoostedModel) -> Result<ShapeFunctions> {
let info = ebm_info(model)?;
let terms = (0..info.terms.len())
.map(|t| info.shape(model, t))
.collect();
Ok(ShapeFunctions {
intercept: f64::from(model.base_scores()[0]) + info.term_means.iter().sum::<f64>(),
terms,
})
}
pub fn term_shape(model: &BoostedModel, term: usize) -> Result<TermShape> {
let info = ebm_info(model)?;
if term >= info.terms.len() {
return Err(HessboostError::incompatible_model(
"term",
format!("the model has {} terms, got {term}", info.terms.len()),
));
}
Ok(info.shape(model, term))
}