use crate::error::{HessboostError, Result};
use serde::{Deserialize, Serialize};
#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize, Default)]
#[serde(rename_all = "lowercase")]
pub enum BoosterKind {
#[default]
GbTree,
Dart,
GbLinear,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize, Default)]
#[serde(rename_all = "lowercase")]
pub enum TreeMethod {
#[default]
Auto,
Exact,
Approx,
Hist,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize, Default)]
#[serde(rename_all = "lowercase")]
pub enum GrowPolicy {
#[default]
DepthWise,
LossGuide,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize, Default)]
#[serde(rename_all = "lowercase")]
pub enum Monotone {
#[default]
None,
Increasing,
Decreasing,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
#[serde(default)]
pub struct TrainingParams {
pub booster: BoosterKind,
pub nthread: usize,
pub seed: u64,
pub objective: String,
pub num_class: usize,
pub base_score: Option<f64>,
pub eval_metric: Vec<String>,
pub tweedie_variance_power: f64,
pub huber_slope: f64,
pub lambdarank_num_pair_per_sample: usize,
pub eta: f64,
pub gamma: f64,
pub max_depth: usize,
pub max_leaves: usize,
pub min_child_weight: f64,
pub max_delta_step: Option<f64>,
pub subsample: f64,
pub colsample_bytree: f64,
pub colsample_bylevel: f64,
pub colsample_bynode: f64,
pub lambda: f64,
pub alpha: f64,
pub scale_pos_weight: f64,
pub tree_method: TreeMethod,
pub grow_policy: GrowPolicy,
pub max_bin: usize,
pub monotone_constraints: Vec<Monotone>,
pub interaction_constraints: Vec<Vec<u32>>,
pub rate_drop: f64,
pub skip_drop: f64,
pub missing: f64,
}
impl Default for TrainingParams {
fn default() -> Self {
TrainingParams {
booster: BoosterKind::GbTree,
nthread: 0,
seed: 0,
objective: "reg:squarederror".to_string(),
num_class: 0,
base_score: None,
eval_metric: Vec::new(),
tweedie_variance_power: 1.5,
huber_slope: 1.0,
lambdarank_num_pair_per_sample: 32,
eta: 0.3,
gamma: 0.0,
max_depth: 6,
max_leaves: 0,
min_child_weight: 1.0,
max_delta_step: None,
subsample: 1.0,
colsample_bytree: 1.0,
colsample_bylevel: 1.0,
colsample_bynode: 1.0,
lambda: 1.0,
alpha: 0.0,
scale_pos_weight: 1.0,
tree_method: TreeMethod::Auto,
grow_policy: GrowPolicy::DepthWise,
max_bin: 256,
monotone_constraints: Vec::new(),
interaction_constraints: Vec::new(),
rate_drop: 0.0,
skip_drop: 0.0,
missing: f64::NAN,
}
}
}
fn ensure(name: &'static str, ok: bool, reason: impl Into<String>) -> Result<()> {
if ok {
Ok(())
} else {
Err(HessboostError::invalid_param(name, reason))
}
}
impl TrainingParams {
pub fn builder() -> TrainingParamsBuilder {
TrainingParamsBuilder {
params: TrainingParams::default(),
}
}
pub fn validate(&self) -> Result<()> {
let unit = |name: &'static str, v: f64| -> Result<()> {
ensure(
name,
v.is_finite() && (0.0..=1.0).contains(&v),
format!("must be in [0, 1], got {v}"),
)
};
let positive = |name: &'static str, v: f64| -> Result<()> {
ensure(
name,
v.is_finite() && v > 0.0,
format!("must be > 0, got {v}"),
)
};
let non_negative = |name: &'static str, v: f64| -> Result<()> {
ensure(
name,
v.is_finite() && v >= 0.0,
format!("must be >= 0, got {v}"),
)
};
positive("eta", self.eta)?;
non_negative("gamma", self.gamma)?;
non_negative("min_child_weight", self.min_child_weight)?;
if let Some(max_delta_step) = self.max_delta_step {
non_negative("max_delta_step", max_delta_step)?;
}
non_negative("lambda", self.lambda)?;
non_negative("alpha", self.alpha)?;
positive("scale_pos_weight", self.scale_pos_weight)?;
unit("subsample", self.subsample)?;
ensure("subsample", self.subsample != 0.0, "must be > 0")?;
unit("colsample_bytree", self.colsample_bytree)?;
unit("colsample_bylevel", self.colsample_bylevel)?;
unit("colsample_bynode", self.colsample_bynode)?;
unit("rate_drop", self.rate_drop)?;
unit("skip_drop", self.skip_drop)?;
if let Some(base_score) = self.base_score {
ensure("base_score", base_score.is_finite(), "must be finite")?;
}
let rho = self.tweedie_variance_power as f32;
ensure(
"tweedie_variance_power",
self.tweedie_variance_power.is_finite() && (1.0f32..2.0).contains(&rho),
format!(
"must be in [1, 2) (as f32), got {}",
self.tweedie_variance_power
),
)?;
positive("huber_slope", self.huber_slope)?;
let slope_sq = (self.huber_slope as f32) * (self.huber_slope as f32);
ensure(
"huber_slope",
slope_sq.is_finite() && slope_sq > 0.0,
format!(
"squared slope must stay positive and finite in f32, got {}",
self.huber_slope
),
)?;
ensure(
"lambdarank_num_pair_per_sample",
self.lambdarank_num_pair_per_sample >= 1,
"must be >= 1",
)?;
ensure(
"max_bin",
self.max_bin >= 2,
format!("must be >= 2, got {}", self.max_bin),
)?;
ensure(
"max_leaves",
!(self.grow_policy == GrowPolicy::LossGuide
&& self.max_leaves == 0
&& self.max_depth == 0),
"lossguide growth needs a bound: set max_leaves or max_depth > 0",
)?;
Ok(())
}
pub fn effective_max_delta_step(&self) -> f64 {
self.max_delta_step
.unwrap_or(if self.objective == "count:poisson" {
0.7
} else {
0.0
})
}
}
#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
pub struct ObjectiveParams {
pub scale_pos_weight: f64,
pub max_delta_step: f64,
pub tweedie_variance_power: f64,
pub huber_slope: f64,
pub lambdarank_num_pair_per_sample: usize,
}
impl ObjectiveParams {
pub fn from_params(p: &TrainingParams) -> Self {
ObjectiveParams {
scale_pos_weight: p.scale_pos_weight,
max_delta_step: p.effective_max_delta_step(),
tweedie_variance_power: p.tweedie_variance_power,
huber_slope: p.huber_slope,
lambdarank_num_pair_per_sample: p.lambdarank_num_pair_per_sample,
}
}
pub fn defaults_for(objective: &str) -> Self {
Self::from_params(
&TrainingParams::builder()
.objective(objective)
.build_unchecked(),
)
}
pub fn training_params(&self, objective: &str, num_class: usize) -> TrainingParamsBuilder {
TrainingParams::builder()
.objective(objective)
.num_class(num_class)
.scale_pos_weight(self.scale_pos_weight)
.max_delta_step(self.max_delta_step)
.tweedie_variance_power(self.tweedie_variance_power)
.huber_slope(self.huber_slope)
.lambdarank_num_pair_per_sample(self.lambdarank_num_pair_per_sample)
}
}
impl Default for ObjectiveParams {
fn default() -> Self {
Self::from_params(&TrainingParams::default())
}
}
#[derive(Debug, Clone)]
pub struct TrainingParamsBuilder {
params: TrainingParams,
}
macro_rules! setter {
($(#[$m:meta])* $name:ident, $ty:ty) => {
$(#[$m])*
#[must_use]
pub fn $name(mut self, v: $ty) -> Self {
self.params.$name = v;
self
}
};
}
impl TrainingParamsBuilder {
setter!( booster, BoosterKind);
setter!( nthread, usize);
setter!( seed, u64);
setter!( num_class, usize);
setter!( eta, f64);
setter!( gamma, f64);
setter!( max_depth, usize);
setter!( max_leaves, usize);
setter!( min_child_weight, f64);
#[must_use]
pub fn max_delta_step(mut self, v: f64) -> Self {
self.params.max_delta_step = Some(v);
self
}
setter!( subsample, f64);
setter!( colsample_bytree, f64);
setter!( colsample_bylevel, f64);
setter!( colsample_bynode, f64);
setter!( lambda, f64);
setter!( alpha, f64);
setter!( scale_pos_weight, f64);
setter!( tree_method, TreeMethod);
setter!( grow_policy, GrowPolicy);
setter!( max_bin, usize);
setter!( rate_drop, f64);
setter!( skip_drop, f64);
setter!( tweedie_variance_power, f64);
setter!( huber_slope, f64);
setter!( lambdarank_num_pair_per_sample, usize);
#[must_use]
pub fn objective(mut self, name: impl Into<String>) -> Self {
self.params.objective = name.into();
self
}
#[must_use]
pub fn base_score(mut self, v: f64) -> Self {
self.params.base_score = Some(v);
self
}
#[must_use]
pub fn eval_metric(mut self, name: impl Into<String>) -> Self {
self.params.eval_metric.push(name.into());
self
}
setter!( monotone_constraints, Vec<Monotone>);
setter!(
interaction_constraints,
Vec<Vec<u32>>
);
pub fn build(self) -> Result<TrainingParams> {
self.params.validate()?;
Ok(self.params)
}
pub fn build_unchecked(self) -> TrainingParams {
self.params
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn defaults_match_xgboost() {
let p = TrainingParams::default();
assert_eq!(p.eta, 0.3);
assert_eq!(p.max_depth, 6);
assert_eq!(p.min_child_weight, 1.0);
assert_eq!(p.lambda, 1.0);
assert_eq!(p.alpha, 0.0);
assert_eq!(p.max_bin, 256);
assert_eq!(p.booster, BoosterKind::GbTree);
assert_eq!(p.grow_policy, GrowPolicy::DepthWise);
assert!(p.base_score.is_none());
assert_eq!(p.tweedie_variance_power, 1.5);
assert_eq!(p.huber_slope, 1.0);
p.validate().unwrap();
}
#[test]
fn builder_chains_and_validates() {
let p = TrainingParams::builder()
.objective("binary:logistic")
.eta(0.1)
.max_depth(4)
.subsample(0.8)
.lambda(2.0)
.build()
.unwrap();
assert_eq!(p.objective, "binary:logistic");
assert_eq!(p.eta, 0.1);
assert_eq!(p.max_depth, 4);
assert_eq!(p.subsample, 0.8);
}
#[test]
fn rejects_bad_params() {
assert!(TrainingParams::builder().eta(0.0).build().is_err());
assert!(TrainingParams::builder().subsample(1.5).build().is_err());
assert!(TrainingParams::builder().lambda(-1.0).build().is_err());
assert!(TrainingParams::builder().max_bin(1).build().is_err());
assert!(
TrainingParams::builder()
.tweedie_variance_power(2.0)
.build()
.is_err()
);
assert!(
TrainingParams::builder()
.tweedie_variance_power(2.0 - f64::EPSILON)
.build()
.is_err()
);
assert!(
TrainingParams::builder()
.tweedie_variance_power(1.0)
.build()
.is_ok()
);
assert!(TrainingParams::builder().huber_slope(0.0).build().is_err());
assert!(TrainingParams::builder().huber_slope(2e19).build().is_err());
assert!(
TrainingParams::builder()
.huber_slope(1e-30)
.build()
.is_err()
);
assert!(
TrainingParams::builder()
.lambdarank_num_pair_per_sample(0)
.build()
.is_err()
);
assert!(
TrainingParams::builder()
.max_delta_step(-1.0)
.build()
.is_err()
);
}
#[test]
fn poisson_delta_step_default_respects_explicit_zero() {
let unset = TrainingParams::builder()
.objective("count:poisson")
.build()
.unwrap();
assert_eq!(unset.effective_max_delta_step(), 0.7);
let zero = TrainingParams::builder()
.objective("count:poisson")
.max_delta_step(0.0)
.build()
.unwrap();
assert_eq!(zero.effective_max_delta_step(), 0.0);
assert_eq!(TrainingParams::default().effective_max_delta_step(), 0.0);
}
#[test]
fn lossguide_requires_bound() {
let r = TrainingParams::builder()
.grow_policy(GrowPolicy::LossGuide)
.max_depth(0)
.max_leaves(0)
.build();
assert!(r.is_err());
TrainingParams::builder()
.grow_policy(GrowPolicy::LossGuide)
.max_leaves(31)
.build()
.unwrap();
}
}