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(¶ms, &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}