Skip to main content

Module ebm

Module ebm 

Source
Expand description

Explainable boosting machines (EBMs): cyclic GA²M boosting, per-term shape functions, and (with Boulevard averaging) confidence bands on them. Opt-in: train with BoosterKind::Ebm, read the shape functions with shape_functions, and, for an Ebm::boulevard model, their bands with EbmInference.

An EBM (Lou, Caruana & Gehrke, KDD 2012; Nori et al., InterpretML, 2019) is a generalized additive model with pairwise interactions (a GA²M): g(E[y | x]) = β + Σ_j f_j(x_j) + Σ_{(j,k)} f_jk(x_j, x_k), each term f learned as a sum of small trees that split only on the term’s features. The trees are ordinary RegTrees, so prediction, SHAP, and every export work as for a gbtree model; EbmInfo records which term each tree belongs to.

§Training

num_boost_round counts EBM rounds; each round grows one tree per term, restricted to the term’s features (the tree shape follows max_depth, max_leaves, grow_policy, min_child_weight, lambda, subsample, …; InterpretML’s defaults are about max_leaves = 3 under loss-guided growth, eta = 0.01, and thousands of rounds). The model counts every tree as one iteration (BoostedModel::num_boost_rounds is the tree count).

  • Classic (the default): cyclic boosting as in InterpretML. Within a round the terms take turns in feature order, each tree fitted to the gradients of the model so far (including the round’s earlier terms) and added with learning rate eta. Any single-output objective works; the shapes are on the margin scale. Categorical features (DMatrix::with_feature_types) get the builders’ native set-membership splits.
  • Outer bags (Ebm::outer_bags = B): each bag boosts every term on its own row sample (Ebm::bag_fraction) and the model averages the bags (each bag’s trees carry 1/B). The bags train in parallel and are combined in bag order. Each tree subsamples its bag’s rows (subsample, by class under BalancedBagging, or by query under QueryBagging).
  • Early stopping (Ebm::early_stopping, an EbmEarlyStopping): InterpretML’s rule. Every bag scores the rows it does not train on after every tree, stops a stage once the last rounds × terms trees failed to beat its best earlier score by its tolerance (relative), and keeps its trees up to its best score, so the bags stop at different rounds and num_boost_round only caps them.
  • Interactions (Ebm::interactions = k): after the main effects, FAST (Lou, Caruana, Gehrke & Hooker, Accurate intelligible models with pairwise interactions, KDD 2013) ranks every pair of features by the best four-quadrant split of the main-effect model’s gradients on their histogram bins (max_bin quantile bins, one bin per category for categorical features, ordered by the category’s mean gradient; rows missing either feature sit out), scored as Σ_q G_q² / (H_q + lambda) − G² / (H + lambda). The top k pairs (ties by feature order) become terms, boosted with the main effects frozen, as InterpretML does.
  • Boulevard (Ebm::boulevard): Fang, Tan, Pipping & Hooker’s inferable EBM (Statistical Inference for Explainable Boosting Machines, AISTATS 2026, Algorithm 1). Round b fits every term’s tree to the same residuals y − ȳ − Σ_t f_t^{(b−1)}(x) (in parallel), centers it on the training rows, t̃ = t − (1/n) Σ_i t(x_i), and averages it into its term, f_t^{(b)} = ((b − 1)/b) f_t^{(b−1)} + (λ/b) t̃ with λ = eta; the model predicts ȳ + ((1 + λ)/λ) Σ_t f_t^{(B)}. Pairs run as a second Boulevard stage on the residuals of the first. See crate::inference for the limit and the bands.

§Shape functions

shape_functions (or term_shape for one term) merges every term’s trees into one piecewise-constant function on the grid their splits cut the term’s features into: the union of the thresholds of a numerical feature, one cell per category a split sends left (plus one for every other category) of a categorical one, and a missing-value cell on each (TermShape, TermAxis), centered to mean zero over the training rows; the intercept collects base_score and the terms’ training means, so intercept + Σ_t shape_t(x) is the model’s margin.

§Refusals

booster = ebm needs one output and no init_model, eval sets, or Trainer::early_stopping_rounds (the terms of one run are fixed; stop each bag with Ebm::early_stopping, flat ebm_early_stopping_rounds, instead); it refuses num_parallel_tree > 1, column sampling, interaction constraints (the terms fix every tree’s features), linear leaves, the reuse penalties, process_type = update, feature weights, and base margins (the shapes and their centering assume the intercept alone). With ebm_boulevard also everything Boulevard inference refuses (non-squared-error objectives, row weights, L1 or clipped leaves, quantized gradients, smoothed leaves, gradient-based sampling), outer bags, early stopping, and base_score. See TrainingParams::validate.

§Deviations from InterpretML

No inner bags, smoothing rounds, or greedy rounds, and early stopping keeps each bag’s best model per stage (InterpretML’s stopping rule and tolerance) without its per-step greedy term selection; pairs use the main effects’ max_bin bins rather than a separate max_interaction_bins, FAST runs once on the bag-averaged main effects rather than per bag, and the shapes are not purified (Lengerich et al., AISTATS 2020): a pair term keeps whatever main-effect part its trees fit.

§Example

use hessboost::config::{BoosterKind, Ebm, GrowPolicy};
use hessboost::ebm::shape_functions;
use hessboost::prelude::*;

let n = 400;
let x: Vec<f32> = (0..n * 2).map(|i| ((i * 37) % 101) as f32 / 101.0).collect();
let y: Vec<f32> = x.chunks(2).map(|r| (6.0 * r[0]).sin() + r[1] * r[1]).collect();
let dtrain = DMatrix::from_dense(&x, n, 2)?.with_labels(&y)?;
let params = TrainingParams::builder()
    .booster(BoosterKind::Ebm(Ebm::default()))
    .eta(0.1)
    .grow_policy(GrowPolicy::LossGuide)
    .max_leaves(3)
    .build()?;
let model = train(&params, &dtrain, 50)?;
let shapes = shape_functions(&model)?;
assert_eq!(shapes.terms.len(), 2);
let margin = shapes.intercept + shapes.terms[0].value(&[0.3])? + shapes.terms[1].value(&[0.5])?;
let direct = model.predict(&DMatrix::from_dense(&[0.3, 0.5], 1, 2)?, Iterations::Best)?;
let direct = *direct.get(0, 0).expect("one row, one output");
assert!((margin - f64::from(direct)).abs() < 1e-4);

Structs§

EbmBoulevard
The settings of a Boulevard EBM its inference reads.
EbmInfo
How a booster = ebm model was trained: the features of every term and which term each tree belongs to, recorded by training (BoostedModel::ebm) and read by shape_functions and EbmInference.
ShapeFunctions
Every term’s shape function and the intercept of an EBM (shape_functions).
TermShape
One term’s shape function: piecewise constant on the grid its trees cut its features into, centered to mean zero over the training rows.

Enums§

TermAxis
The cells of one feature of a term (TermShape::axes). Either way the last cell holds missing values.

Functions§

shape_functions
The shape functions and intercept of model, an EBM (BoosterKind::Ebm): every term’s trees merged into one piecewise-constant function of its features (see TermShape), missing values included. intercept + Σ_t shape_t(x) is the model’s margin (up to f32 rounding of the tree sum).
term_shape
The shape function of term term of the EBM model: element term of shape_functions’ terms, built without the other terms’ grids.