Skip to main content

hessboost/
lib.rs

1//! # hessboost
2//!
3//! Fast, deterministic gradient boosting in Rust (with Python bindings).
4//! hessboost provides multi-core tree building with runtime-detected NEON and AVX2
5//! SIMD, strict parameter validation, reproducible models on any thread count, and
6//! stable model storage. It supports modern extensions like conformal prediction,
7//! explainable boosting machines (EBMs), distributional modeling, and tree-based
8//! diffusion, alongside bidirectional XGBoost JSON/UBJSON model interchange.
9//!
10//! ## Quick start
11//!
12//! Build a [`DMatrix`], set [`TrainingParams`] with its builder, train with
13//! [`train`] (or [`Trainer`] for eval sets, early stopping, custom hooks,
14//! and continued training), then
15//! [`predict`](model::BoostedModel::predict):
16//!
17//! ```
18//! use hessboost::prelude::*;
19//!
20//! # fn main() -> Result<()> {
21//! // 6 rows × 2 features, row-major, plus a label per row.
22//! let x = [0.0, 0.0,  1.0, 0.0,  0.0, 1.0,  1.0, 1.0,  0.5, 0.5,  0.2, 0.9];
23//! let y = [0.0,       1.0,       1.0,       0.0,       0.5,       0.7];
24//! let dtrain = DMatrix::from_dense(&x, 6, 2)?.with_labels(&y)?;
25//!
26//! let params = TrainingParams::builder()
27//!     .objective(Objective::SquaredError(RegLoss::default()))
28//!     .tree_method(TreeMethod::Hist)
29//!     .max_depth(3)
30//!     .eta(0.1)
31//!     .build()?;
32//!
33//! let model = train(&params, &dtrain, 50)?;
34//! let preds = model.predict(&dtrain, Iterations::Best)?;
35//! assert_eq!((preds.n_rows(), preds.width()), (6, 1)); // `[row][output]`
36//!
37//! model.save("model.bin", ModelFormat::Binary)?; // native format
38//! # std::fs::remove_file("model.bin").ok();
39//! # Ok(())
40//! # }
41//! ```
42//!
43//! [`prelude`] holds only this workflow's items (including
44//! [`TreeMethod`](config::TreeMethod) and [`Objective`](objective::Objective));
45//! everything else is imported from its module.
46//!
47//! ## Modules
48//!
49//! - [`config`]: [`TrainingParams`], builder, parameter enums.
50//! - [`data`]: [`DMatrix`], [`MetaInfo`](data::MetaInfo), feature types,
51//!   CSV/libsvm loaders, [`data::target_stats`].
52//! - [`training`]: [`train`], [`Trainer`], [`cv`](training::cv),
53//!   [`training::budget`], [`training::online`].
54//! - [`model`]: [`BoostedModel`] (prediction, SHAP, importance, slicing,
55//!   native and XGBoost JSON/UBJSON, LightGBM text import);
56//!   [`model::compact`], [`model::uncertainty`].
57//! - [`objective`]: [`Objective`](objective::Objective) and its parameter
58//!   types, the `Loss` trait, `CustomLoss`, [`objective::distributional`]
59//!   (`dist:*` objectives).
60//! - [`metric`]: [`EvalMetric`](metric::EvalMetric) (the built-in metrics),
61//!   the `Metric` trait, `CustomMetric`.
62//! - [`conformal`]: split-conformal and conformalized-quantile intervals.
63//! - [`inference`]: Boulevard boosting's confidence and prediction intervals
64//!   for `f(x)`, and a Boulevard EBM's shape-function bands.
65//! - [`ebm`]: explainable boosting machines' terms and shape functions.
66//! - [`diffusion`]: conditional diffusion and flow matching with GBDT score
67//!   models, sampling a nonparametric `p(y | x)`.
68//! - [`tree`]: [`RegTree`](tree::RegTree) and nodes, for model inspection.
69//! - [`error`]: `HessboostError` and `Result`.
70//!
71//! ## What's here
72//!
73//! - **Boosters:** [`BoosterKind::GbTree`](config::BoosterKind::GbTree),
74//!   [`BoosterKind::Dart`](config::BoosterKind::Dart),
75//!   [`BoosterKind::GbLinear`](config::BoosterKind::GbLinear), boosted random
76//!   forests (`num_parallel_tree`),
77//!   [`BoosterKind::Boulevard`](config::BoosterKind::Boulevard), and
78//!   [`BoosterKind::Ebm`](config::BoosterKind::Ebm).
79//! - **Lifecycle:** continued training and `process_type=update` refresh
80//!   ([`Trainer::init_model`](training::Trainer::init_model)), a per-round
81//!   hook for progress, custom stopping, and cancellation
82//!   ([`Trainer::on_round`](training::Trainer::on_round)), slicing
83//!   ([`BoostedModel::slice`](model::BoostedModel::slice)), `iteration_range`
84//!   prediction as Rust ranges (every prediction method's
85//!   [`Iterations`](model::Iterations) argument).
86//! - **Tree methods:** `exact`, `hist`, `approx`; `depthwise`/`lossguide`
87//!   growth; uniform or `gradient_based` row sampling; column sampling with
88//!   optional per-feature weights
89//!   ([`DMatrix::with_feature_weights`](data::DMatrix::with_feature_weights)).
90//! - **Objectives** ([`Objective`](objective::Objective), each with its
91//!   parameters): regression (squared, squared-log, pseudo-Huber, smoothed
92//!   absolute, quantile/expectile lists), binary (logistic, logitraw, hinge)
93//!   and multiclass, counts, LambdaMART and XE-NDCG ranking, survival (`survival:cox`,
94//!   `survival:aft` on censored bounds), plus custom losses
95//!   ([`Objective::Custom`](objective::Objective::Custom), e.g. a
96//!   [`CustomLoss`](objective::CustomLoss)).
97//! - **Multi-output:** label matrices
98//!   ([`DMatrix::with_label_matrix`](data::DMatrix::with_label_matrix)), one
99//!   tree per output or vector-leaf trees
100//!   ([`MultiStrategy::MultiOutputTree`](config::MultiStrategy::MultiOutputTree)).
101//! - **Metrics** ([`EvalMetric`](metric::EvalMetric), each with its own
102//!   parameters): rmse, rmsle, mae, mape, mphe, logloss, error, auc, aucpr,
103//!   mlogloss, merror, poisson/gamma/tweedie-nloglik, ndcg, map, pre,
104//!   quantile, expectile, cox/aft-nloglik, interval-regression-accuracy, plus
105//!   a custom hook ([`Trainer::custom_metric`](training::Trainer::custom_metric),
106//!   reported after built-in metrics). XGBoost-compatible parameter dictionaries parse through
107//!   [`TrainingParams::from_xgboost`](config::TrainingParams::from_xgboost):
108//!   `@k` ranking cutoffs and `@rho` on tweedie-nloglik, other suffixes
109//!   refused.
110//! - **Modeling:** monotone and interaction constraints, native categorical
111//!   splits, early stopping, feature importance, QuadratureTreeSHAP values
112//!   and interactions
113//!   ([`predict_contribs`](model::BoostedModel::predict_contribs) /
114//!   [`predict_interactions`](model::BoostedModel::predict_interactions)).
115//! - **I/O:** libsvm/CSV loaders, native binary + JSON, XGBoost JSON and
116//!   UBJSON import/export ([XGBoost interchange](model#xgboost-interchange)),
117//!   LightGBM 4.x text model import
118//!   ([`ModelFormat::LightgbmText`](model::ModelFormat::LightgbmText); see
119//!   [LightGBM import](model#lightgbm-import)), all through one
120//!   [`ModelFormat`](model::ModelFormat) with byte-level
121//!   [`detect`](model::ModelFormat::detect)ion; models compiled into the
122//!   binary with [`EmbeddedModel`](model::EmbeddedModel).
123//! - **Validation:** cross-validation ([`cv`](training::cv)), custom,
124//!   forward-chaining (time-ordered, purged by a row gap), or purged forward
125//!   (timestamped rows, purged by each label window,
126//!   [`Fold::purged_forward`](training::Fold::purged_forward))
127//!   [`Fold`](training::Fold)s with fold-mean early stopping
128//!   ([`CrossValidation`](training::CrossValidation)); whole-query folds of
129//!   ranking data; ordered target statistics fitted inside each fold
130//!   ([`CrossValidation::target_stats`](training::CrossValidation::target_stats)).
131//! - **Advanced & experimental methods (opt-in):**
132//!   - split-conformal and conformalized-quantile intervals with
133//!     finite-sample marginal coverage ([`conformal`]);
134//!   - Boulevard boosting (Zhou & Hooker, JMLR 2022) and its dropout
135//!     (BRAT-D) and parallel (BRAT-P) variants (Fang, Tan & Hooker, NeurIPS
136//!     2025) with CLT-based confidence intervals for `f(x)`, prediction and
137//!     reproduction intervals, an
138//!     honest leaf refit, and exact or Nyström variance
139//!     ([`BoosterKind::Boulevard`](config::BoosterKind::Boulevard),
140//!     [`inference`]);
141//!   - explainable boosting machines (GA²M: cyclic per-feature boosting,
142//!     early-stopped outer bags, FAST pair terms, numerical and categorical
143//!     terms; Lou et al., KDD 2012/2013, InterpretML)
144//!     with per-term shape functions, and their Boulevard variant (Fang, Tan,
145//!     Pipping & Hooker, AISTATS 2026) with confidence bands on every shape
146//!     ([`BoosterKind::Ebm`](config::BoosterKind::Ebm), [`ebm`],
147//!     [`EbmInference`](inference::EbmInference));
148//!   - CatBoost-style ordered target statistics ([`data::target_stats`]);
149//!   - LightGBM options `extra_trees`, `path_smooth`, `linear_tree` leaves
150//!     ([`TrainingParams::extra_trees`](config::TrainingParams::extra_trees),
151//!     [`path_smooth`](config::TrainingParams::path_smooth),
152//!     [`linear_tree`](config::TrainingParams::linear_tree),
153//!     [`LinearLeaves`](tree::LinearLeaves));
154//!   - LightGBM class-balanced bagging for binary classification
155//!     ([`BalancedBagging`](config::BalancedBagging): `pos_bagging_fraction`,
156//!     `neg_bagging_fraction`), in place of `subsample`;
157//!   - LightGBM XE-NDCG ranking
158//!     ([`Objective::RankXendcg`](objective::Objective::RankXendcg); its keyed
159//!     per-round draws differ from LightGBM's random stream) and query-level
160//!     bagging ([`QueryBagging`](config::QueryBagging), `bagging_by_query`);
161//!   - CatBoost-style symmetric trees
162//!     ([`GrowPolicy::Symmetric`](config::GrowPolicy::Symmetric)), routed by
163//!     bit pattern in batch prediction;
164//!   - *Boosted Trees on a Diet* reuse penalties
165//!     (`toad_penalty_feature`, `toad_penalty_threshold`) and a bit-packed
166//!     layout with bit-identical margins ([`model::compact`]);
167//!   - LightGBM-style quantized gradients
168//!     ([`QuantizedGrad`](config::QuantizedGrad), `use_quantized_grad`);
169//!   - PerpetualBooster-style budget training: one `budget` instead of
170//!     `eta`/depth/rounds ([`training::budget`]);
171//!   - in-place row addition and deletion (incremental learning and machine
172//!     unlearning) for trained hist models, exact or approximate
173//!     ([`training::online`]);
174//!   - distributional boosting (NGBoost / XGBoostLSS style): `dist:normal`,
175//!     `dist:lognormal`, `dist:gamma`, `dist:poisson`, `dist:negbinomial`
176//!     per-row distributions
177//!     ([`predict_distribution`](model::BoostedModel::predict_distribution),
178//!     [`objective::distributional`]), scored by `nll` / `crps`;
179//!   - CatBoost's Stochastic Gradient Langevin Boosting and model shrinkage
180//!     ([`langevin`](config::TrainingParams::langevin),
181//!     [`model_shrink`](config::TrainingParams::model_shrink),
182//!     [`posterior_sampling`](config::TrainingParams::posterior_sampling))
183//!     with virtual ensembles: knowledge, data, and total uncertainty from
184//!     one model's exactly rebuilt truncations
185//!     ([`predict_uncertainty`](model::BoostedModel::predict_uncertainty),
186//!     [`model::uncertainty`]);
187//!   - nonparametric `p(y | x)` by tree-based conditional diffusion
188//!     (Treeffuser) and flow matching (DiffGBM) for scalar or vector
189//!     labels, sampled deterministically ([`diffusion`]);
190//!   - ForestFlow / ForestDiffusion tabular generation and imputation
191//!     ([`diffusion::forest`]);
192//!   - native Metal on macOS 10.15+ (`metal` feature): bit-identical GPU
193//!     prediction ([`to_gpu`](model::BoostedModel::to_gpu), ~3x faster at
194//!     scale) and bit-identical GPU histograms
195//!     ([`device`](config::TrainingParams::device) = `metal`; exact integer
196//!     sums, CPU fallback outside their exact domain). Documented only in
197//!     macOS builds with the feature (`cargo doc --features metal`);
198//!     elsewhere [`backend::metal`] is a stub.
199//!
200//! `examples/` has one program per topic (`train_regression`,
201//! `binary_classification`, `multiclass`, `ranking`, `rank_xendcg`, `shap`,
202//! `model_io`, `custom_objective`, `constraints`, `conformal`,
203//! `boulevard_inference`, `ebm`, `compact_model`, `distributional`,
204//! `virtual_ensembles`, `tree_diffusion`, `forest_flow`, `budget`,
205//! `balanced_bagging`, `online_update`, `ordered_target_stats`, `pfn_boost`,
206//! `metal` with `--features metal` on macOS). Run one with
207//! `cargo run --release --example binary_classification`.
208//!
209//! ## Compatibility notes
210//!
211//! While hessboost is a standalone library, it offers extensive compatibility with
212//! XGBoost configurations and models:
213//! [`TrainingParams::from_xgboost`](config::TrainingParams::from_xgboost)
214//! reads XGBoost `params` dictionaries (keys, aliases, and value spellings), and
215//! unsupported settings are refused. Deterministic configurations reproduce XGBoost
216//! predictions within `1e-4` (quantile cuts bit for bit), and imported XGBoost models
217//! predict and explain identically. RNG-driven options (subsampling,
218//! forests, DART) match in quality only — the random streams differ.
219//! ### Not implemented
220//!
221//! - Distributed and external-memory training; GPU training outside macOS.
222//! - XGBoost options available at one setting only (so they are not
223//!   [`TrainingParams`] fields; `from_xgboost` accepts exactly that
224//!   setting): gblinear uses `updater = coord_descent`
225//!   with `feature_selector = cyclic`; LambdaMART uses
226//!   `lambdarank_pair_method = topk` (no `lambdarank_unbiased` or
227//!   `ndcg_exp_gain`); DART has no `sample_type` or `normalize_type` (it
228//!   samples uniformly and normalizes by `tree`); categorical splits use
229//!   XGBoost's defaults
230//!   `max_cat_to_onehot = 4` and `max_cat_threshold = 64`.
231//! - The metrics `gamma-deviance`, `error@t` (XGBoost's classification
232//!   threshold suffix), and the `-` variants of the ranking metrics
233//!   (`ndcg-`, `ndcg@k-`, `map-`, `map@k-`); these names are refused.
234//! - XGBoost import and export of gblinear models.
235//! - `scale_pos_weight = 0`: XGBoost's bound is `>= 0`; hessboost's
236//!   [`RegLoss::new`](objective::RegLoss::new) needs a positive weight, so
237//!   configurations and XGBoost files with `0` are refused.
238//!
239//! [`DMatrix`]: data::DMatrix
240//! [`TrainingParams`]: config::TrainingParams
241//! [`BoostedModel`]: model::BoostedModel
242//! [`train`]: training::train
243//! [`Trainer`]: training::Trainer
244#![forbid(unsafe_op_in_unsafe_fn)]
245#![warn(missing_docs)]
246
247pub mod backend;
248mod check;
249pub mod config;
250pub mod conformal;
251pub mod data;
252pub mod diffusion;
253pub mod ebm;
254pub mod error;
255pub mod inference;
256pub mod metric;
257pub mod model;
258pub mod objective;
259mod rng;
260mod simd;
261#[cfg(test)]
262mod test_support;
263pub mod training;
264pub mod tree;
265/// `1e-6` in `f64` arithmetic, where the crate compares against XGBoost's
266/// `kRtEps` in double precision (the `f64` literal, not [`K_RT_EPS_F32`]
267/// widened).
268pub(crate) const K_RT_EPS: f64 = 1e-6;
269/// XGBoost's `kRtEps` (`1e-6f`): the minimum gain improvement a split must
270/// beat, and the floor of sampling weights and near-zero sums.
271pub(crate) const K_RT_EPS_F32: f32 = 1e-6;
272/// The train-and-predict workflow in one import: `use hessboost::prelude::*;`.
273///
274/// Holds the data container, the parameters, the training entry points, the
275/// model, the error types, and the types their everyday methods take:
276/// [`TreeMethod`](config::TreeMethod) (for
277/// [`TrainingParamsBuilder::tree_method`](config::TrainingParamsBuilder::tree_method)),
278/// [`ImportanceType`](model::ImportanceType) (for
279/// [`BoostedModel::feature_importance`](model::BoostedModel::feature_importance)),
280/// [`Objective`](objective::Objective) (for
281/// [`TrainingParamsBuilder::objective`](config::TrainingParamsBuilder::objective))
282/// with [`RegLoss`](objective::RegLoss) (the parameter of its default
283/// `reg:squarederror`), and [`EvalMetric`](metric::EvalMetric) (for
284/// [`TrainingParamsBuilder::eval_metric`](config::TrainingParamsBuilder::eval_metric)).
285/// Everything else (the other parameter enums, the other objectives' and
286/// metrics' parameters, conformal intervals, ...) is imported from its module.
287pub mod prelude {
288    pub use crate::config::{TrainingParams, TreeMethod};
289    pub use crate::data::DMatrix;
290    pub use crate::error::{HessboostError, Result};
291    pub use crate::metric::EvalMetric;
292    pub use crate::model::{BoostedModel, ImportanceType, Iterations, ModelFormat};
293    pub use crate::objective::{Objective, RegLoss};
294    pub use crate::training::{Trainer, train};
295}
296/// Implementation details the crate's own benchmarks and parity tests
297/// drive directly (histogram construction, tree growth, quantile cuts). Not
298/// part of the public API: hidden from the docs and changed without notice.
299#[doc(hidden)]
300pub mod internals {
301    pub use crate::data::ghist::GHistIndex;
302    pub use crate::data::quantile::HistCuts;
303    pub use crate::tree::builder::HistTreeBuilder;
304    pub use crate::tree::hist::{CpuBackend, HistogramBackend, zeroed};
305    pub use crate::tree::sampler::ColumnSampler;
306}