Skip to main content

hessboost/
lib.rs

1//! # hessboost
2//!
3//! A faithful, fast, pure-Rust reimplementation of
4//! [XGBoost](https://github.com/dmlc/xgboost) gradient boosting with no C/C++
5//! dependency and no FFI.
6//!
7//! ## Quick start
8//!
9//! Build a [`DMatrix`], configure [`TrainingParams`] with a builder, call
10//! [`train`], then [`predict`](prelude::BoostedModel::predict):
11//!
12//! ```
13//! use hessboost::prelude::*;
14//!
15//! # fn main() -> Result<()> {
16//! // 6 rows × 2 features, row-major, plus a label per row.
17//! 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];
18//! let y = [0.0,       1.0,       1.0,       0.0,       0.5,       0.7];
19//! let dtrain = DMatrix::from_dense(&x, 6, 2)?.with_labels(&y)?;
20//!
21//! let params = TrainingParams::builder()
22//!     .objective("reg:squarederror") // XGBoost-compatible names
23//!     .tree_method(TreeMethod::Hist)
24//!     .max_depth(3)
25//!     .eta(0.1)
26//!     .build()?;
27//!
28//! let model = train(&params, &dtrain, 50)?;
29//! let preds = model.predict(&dtrain)?;
30//! assert_eq!(preds.len(), 6);
31//!
32//! model.save_binary("model.bin")?;      // native format
33//! # std::fs::remove_file("model.bin").ok();
34//! # Ok(())
35//! # }
36//! ```
37//!
38//! ## What's here
39//!
40//! - **Boosters:** `gbtree`, `dart`, `gblinear`.
41//! - **Tree methods:** `exact`, `hist`, and `approx`, with `depthwise` or
42//!   `lossguide` growth.
43//! - **Objectives:** regression, binary/multiclass classification, count
44//!   (poisson/gamma/tweedie), learning-to-rank (LambdaMART), and a custom hook
45//!   ([`train_with_objective`]).
46//! - **Metrics:** rmse, mae, logloss, error, auc, aucpr, mlogloss, merror,
47//!   ndcg/map, nloglik, and a custom hook ([`train_with_custom_metric`]).
48//! - **Modeling:** monotone & interaction constraints, native categorical
49//!   splits, early stopping, feature importance, TreeSHAP contributions and
50//!   interaction values ([`BoostedModel::predict_contribs`] /
51//!   [`predict_interactions`](prelude::BoostedModel::predict_interactions)).
52//! - **I/O:** libsvm/CSV loaders, native binary + JSON model I/O, and
53//!   XGBoost-format JSON model import/export ([`crate::model`]).
54//! - **Validation:** cross-validation ([`cv`]).
55//!
56//! ## Where to look
57//!
58//! - Entry points: [`train`], [`train_with_eval`], [`train_with_objective`],
59//!   [`train_with_custom_metric`], [`cv`].
60//! - Core types: [`DMatrix`] (data), [`TrainingParams`] (config, mirrors
61//!   XGBoost parameter names), [`BoostedModel`] (trained model).
62//! - Runnable examples in the crate's `examples/` directory (e.g.
63//!   `binary_classification`, `multiclass`, `ranking`, `shap`, `model_io`,
64//!   `custom_objective`, `constraints`). Run one with
65//!   `cargo run --release --example binary_classification`.
66//!
67//! ## Compatibility notes
68//!
69//! Objective, metric, and parameter names mirror XGBoost, so configurations
70//! transfer directly. Predictions match XGBoost's *model quality* (parity is
71//! CI-tested) but are not bit-identical. The two histogram implementations pick
72//! slightly different split points.
73//!
74//! [`DMatrix`]: prelude::DMatrix
75//! [`TrainingParams`]: prelude::TrainingParams
76//! [`BoostedModel`]: prelude::BoostedModel
77//! [`BoostedModel::predict_contribs`]: prelude::BoostedModel::predict_contribs
78//! [`train`]: prelude::train
79//! [`train_with_eval`]: prelude::train_with_eval
80//! [`train_with_objective`]: prelude::train_with_objective
81//! [`train_with_custom_metric`]: prelude::train_with_custom_metric
82//! [`cv`]: prelude::cv
83#![forbid(unsafe_op_in_unsafe_fn)]
84#![warn(missing_docs)]
85
86pub mod booster;
87pub mod config;
88pub mod data;
89pub mod error;
90pub mod learner;
91pub mod metric;
92pub mod model;
93pub mod objective;
94mod simd;
95pub mod tree;
96
97/// Commonly used imports include `use hessboost::prelude::*;`.
98///
99/// Pulls in the data container, configuration, training entry points, the model
100/// type, and the objective/metric hooks. This provides everything needed for the
101/// typical train to predict workflow.
102pub mod prelude {
103    pub use crate::config::{BoosterKind, GrowPolicy, Monotone, TrainingParams, TreeMethod};
104    pub use crate::data::{CsvOptions, DMatrix, FeatureType};
105    pub use crate::error::{HessboostError, Result};
106    pub use crate::learner::{
107        BoostedModel, CvResult, ImportanceType, TrainResult, cv, train, train_with_custom_metric,
108        train_with_eval, train_with_objective,
109    };
110    pub use crate::metric::{CustomMetric, Metric};
111    pub use crate::objective::{CustomObjective, GradPair, Objective};
112}