use crate::check::ensure;
use crate::config::TrainingParams;
use crate::data::DMatrix;
use crate::data::ghist::GHistIndex;
use crate::data::quantile::HistCuts;
use crate::error::{HessboostError, Result};
use crate::model::BoostedModel;
use crate::objective::{GradPair, Objective};
use crate::training::multi_output::reject_split_gradient;
use crate::training::train::{initial_intercepts, new_model, with_thread_pool};
use crate::training::validate::{
reject_feature_weights, validate_datasets, validate_trained_model,
};
use crate::tree::builder::budget::{
ChildRecord, GENERALIZATION_THRESHOLD_RELAXED, GrowConfig, N_FOLDS, TreeStopper,
fold_weight_spread, grow,
};
pub const DEFAULT_BUDGET: f64 = 0.5;
pub const MAX_BUDGET: f64 = 5.0;
const STOPPING_ROUNDS: usize = 3;
const ITER_LIMIT: usize = 1000;
#[derive(Debug, Clone, Copy, PartialEq)]
#[non_exhaustive]
pub struct BudgetConfig {
pub budget: f64,
pub iteration_limit: Option<usize>,
pub stopping_rounds: Option<usize>,
}
impl Default for BudgetConfig {
fn default() -> Self {
BudgetConfig::new(DEFAULT_BUDGET)
}
}
impl BudgetConfig {
pub fn new(budget: f64) -> Self {
BudgetConfig {
budget,
iteration_limit: None,
stopping_rounds: None,
}
}
#[must_use]
pub fn iteration_limit(mut self, limit: usize) -> Self {
self.iteration_limit = Some(limit);
self
}
#[must_use]
pub fn stopping_rounds(mut self, rounds: usize) -> Self {
self.stopping_rounds = Some(rounds);
self
}
pub fn validate(&self) -> Result<()> {
ensure(
"budget",
self.budget > 0.0 && self.budget < MAX_BUDGET,
format!("must be in (0, {MAX_BUDGET}), got {}", self.budget),
)?;
ensure(
"iteration_limit",
self.iteration_limit != Some(0),
"must be at least 1",
)?;
ensure(
"stopping_rounds",
self.stopping_rounds != Some(0),
"must be at least 1",
)?;
Ok(())
}
pub fn eta(&self) -> f64 {
let b = self.budget.max(0.0);
let power = if b <= 1.0 { b } else { 1.0 + 0.65 * (b - 1.0) };
10f64.powf(-power)
}
fn scale(&self, exponent: f64, max: f64) -> f64 {
10f64
.powf((self.budget - 1.0).max(0.0) * exponent)
.clamp(1.0, max)
}
pub(crate) fn effective_stopping_rounds(&self) -> usize {
self.stopping_rounds
.unwrap_or_else(|| (STOPPING_ROUNDS as f64 * self.scale(0.5, 6.0)).ceil() as usize)
}
pub(crate) fn effective_iteration_limit(&self) -> usize {
let derived = (ITER_LIMIT as f64 * self.scale(0.35, 4.0)).round() as usize;
self.iteration_limit
.map_or(derived, |limit| limit.min(derived))
}
fn base_target(&self, loss_avg: f64) -> f64 {
let u = self.budget.max(0.1);
let n = 10.0 / u;
let c = (n - 2.0) / (n * (n - 1.0));
c * 10f64.powf(-u.min(3.0)) * loss_avg.max(f64::from(f32::EPSILON))
}
fn target(&self, initial_loss: f64, previous_loss: f64) -> f64 {
if self.budget <= 0.2 {
return self.base_target(initial_loss);
}
let alpha = (1.0 - 0.1 * self.budget).clamp(0.85, 0.995);
self.base_target(alpha * initial_loss + (1.0 - alpha) * previous_loss) * 1.2
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
#[non_exhaustive]
pub enum BudgetStop {
RootUnsplittable,
WeakTrees,
NoImprovement,
IterationLimit,
NonFiniteLoss,
}
#[derive(Debug)]
#[non_exhaustive]
pub struct BudgetResult {
pub model: BoostedModel,
pub eta: f64,
pub stop: BudgetStop,
}
pub fn train_with_budget(
params: &TrainingParams,
dtrain: &DMatrix,
config: &BudgetConfig,
) -> Result<BudgetResult> {
let result = with_thread_pool(params, || train_budget_inner(params, dtrain, config))?;
validate_trained_model(&result.model)?;
Ok(result)
}
fn reject_tuned_params(params: &TrainingParams) -> Result<()> {
let mut reference = TrainingParams {
objective: params.objective.clone(),
..TrainingParams::default()
};
if matches!(params.objective, Objective::Poisson) {
reference.max_delta_step = params.max_delta_step;
}
reference.base_score = params.base_score;
reference.max_bin = params.max_bin;
reference.nthread = params.nthread;
params.refuse_changes_from(
&reference,
"budget",
"budget mode derives the learning rate, tree shape, and round count itself",
)
}
fn train_budget_inner(
params: &TrainingParams,
dtrain: &DMatrix,
config: &BudgetConfig,
) -> Result<BudgetResult> {
config.validate()?;
params.validate()?;
reject_feature_weights(dtrain, "budget mode does not sample columns")?;
let objective = params.loss(dtrain.n_targets())?;
let n_out = objective.n_outputs();
let loss_fn = match objective.pointwise_loss() {
Some(loss) if n_out == 1 => loss,
_ => {
return Err(HessboostError::invalid_param(
"objective",
format!(
"`{}` is not supported by budget mode (it needs a single-output \
objective with a pointwise loss)",
objective.name()
),
));
}
};
reject_tuned_params(params)?;
let Some(labels) = dtrain.labels() else {
return Err(HessboostError::EmptyDataset(
"train_with_budget: dtrain has no labels",
));
};
let n = dtrain.n_rows();
validate_datasets(objective.as_ref(), dtrain, &[])?;
let info = dtrain.info();
let base_margins = initial_intercepts(params, objective.as_ref(), &info, n_out)?;
let mut model = new_model(params, objective.as_ref(), dtrain, base_margins);
let ghist = GHistIndex::from_dmatrix(dtrain, HistCuts::from_dmatrix(dtrain, params.max_bin));
let weights = dtrain.weights();
let weight_of = |r: usize| weights.map_or(1.0, |w| f64::from(w[r]));
let row_loss = |r: usize, margin: f32| weight_of(r) * loss_fn(margin, labels[r]);
let mut margins = model.initial_margins(dtrain);
let mut loss: Vec<f64> = (0..n).map(|r| row_loss(r, margins[r])).collect();
let average = |loss: &[f64]| loss.iter().sum::<f64>() / n.max(1) as f64;
let eta = config.eta();
let stopping_rounds = config.effective_stopping_rounds();
let regression_like = matches!(
params.objective,
Objective::SquaredError(_) | Objective::PseudoHuber(_)
);
let initial_loss = average(&loss);
let mut previous_loss = initial_loss;
let mut best_loss = initial_loss;
let mut weak_rounds = 0usize;
let mut untargeted_rounds = 0usize;
let mut no_improvement = 0usize;
let mut gpair = vec![GradPair::default(); n];
let mut stop = BudgetStop::IterationLimit;
for round in 0..config.effective_iteration_limit() {
let target = (untargeted_rounds <= stopping_rounds.saturating_add(1))
.then(|| config.target(initial_loss, previous_loss));
objective.gradient_info(&margins, &info, &mut gpair);
reject_split_gradient(objective.as_ref(), round, &gpair)?;
let row_decrement = |r: u32, delta: f32| {
let r = r as usize;
loss[r] - row_loss(r, margins[r] + delta)
};
let grown = grow(
&ghist,
&gpair,
&GrowConfig {
eta: eta as f32,
target_loss_decrement: target,
row_decrement: &row_decrement,
max_delta_step: params.effective_max_delta_step(),
},
)?;
grown.apply(&mut margins);
let n_nodes = grown.tree.num_nodes();
let generalization = tree_generalization(&grown.children, regression_like);
let mut stop_now = false;
if n_nodes < 5
&& generalization < GENERALIZATION_THRESHOLD_RELAXED
&& grown.stopper != TreeStopper::StepSize
{
weak_rounds += 1;
stop_now = n_nodes == 1;
}
if grown.stopper == TreeStopper::StepSize {
untargeted_rounds = 0;
} else {
untargeted_rounds += 1;
}
for (r, l) in loss.iter_mut().enumerate() {
*l = row_loss(r, margins[r]);
}
let current_loss = average(&loss);
if !current_loss.is_finite() {
stop = BudgetStop::NonFiniteLoss;
break;
}
previous_loss = current_loss;
if current_loss < best_loss {
best_loss = current_loss;
no_improvement = 0;
} else {
no_improvement += 1;
}
model.push_tree_weighted(grown.tree, 1.0);
if stop_now {
stop = BudgetStop::RootUnsplittable;
break;
}
if weak_rounds >= stopping_rounds {
stop = BudgetStop::WeakTrees;
break;
}
if no_improvement >= stopping_rounds {
stop = BudgetStop::NoImprovement;
break;
}
}
Ok(BudgetResult { model, eta, stop })
}
fn fold_weight_reliability(weights: &[f64; N_FOLDS]) -> f64 {
fold_weight_spread(weights).map_or(1.0, |(mean_abs, std_dev)| {
let positive = weights.iter().filter(|&&w| w >= 0.0).count() as f64 / N_FOLDS as f64;
(positive.max(1.0 - positive) / (1.0 + std_dev / mean_abs)).clamp(0.5, 1.0)
})
}
fn tree_generalization(children: &[ChildRecord], regression_like: bool) -> f64 {
let mut best = 0.0f64;
let (mut weighted, mut total) = (0.0, 0.0);
for child in children {
let stability = fold_weight_reliability(&child.fold_weights);
let node_score = child.generalization * (0.99 + 0.01 * stability);
best = best.max(node_score);
let node_weight = (child.count.max(1) as f64).sqrt() * stability;
weighted += node_score.clamp(0.95, 1.05) * node_weight;
total += node_weight;
}
match (regression_like, total > 0.0) {
(true, true) => weighted / total,
(true, false) => 0.0,
(false, _) => best,
}
}