use super::dart::{round_rng, round_salt};
use super::round::gather_output;
use super::row_sampling::{RowMeta, gradient_sampling};
use crate::config::{Device, GrowPolicy, TrainingParams, TreeMethod};
use crate::data::ghist::GHistIndex;
use crate::data::quantile::HistCuts;
use crate::data::{DMatrix, MetaInfo};
use crate::error::{HessboostError, Result};
use crate::objective::{GradPair, Loss};
use crate::training::sampling::gradient_based_sample;
use crate::training::sglb::Langevin;
use crate::tree::RegTree;
use crate::tree::builder::{
ExactTreeBuilder, HistTreeBuilder, LeafRows, SortedColumns, check_symmetric_input,
};
use crate::tree::hist::{CpuBackend, HistogramBackend};
use crate::tree::reuse::ReuseSet;
use crate::tree::sampler::ColumnSampler;
use std::sync::OnceLock;
#[derive(Clone, Copy)]
pub(super) struct TrainContext<'a> {
pub(super) params: &'a TrainingParams,
pub(super) dtrain: &'a DMatrix,
pub(super) info: &'a MetaInfo<'a>,
pub(super) objective: &'a dyn Loss,
pub(super) langevin: Option<&'a Langevin>,
pub(super) rows: RowMeta<'a>,
}
#[derive(Clone, Copy)]
pub(super) struct TreeSample<'a> {
pub(super) gpair: &'a [GradPair],
pub(super) rows: &'a [u32],
pub(super) forest_index: Option<&'a OnceLock<GHistIndex>>,
}
pub(super) enum Prepared {
Exact(SortedColumns),
Hist {
index: GHistIndex,
backend: Box<dyn HistogramBackend>,
rows_route_like_trees: bool,
},
Approx {
const_hess: bool,
cached: OnceLock<GHistIndex>,
},
}
impl Prepared {
pub(super) fn build_tree(
&self,
run: &TrainContext,
sample: TreeSample,
sampler: &mut ColumnSampler,
reuse: Option<&mut ReuseSet>,
rounding_seed: u64,
capture_rows: bool,
) -> (RegTree, Vec<LeafRows>) {
let TrainContext { params, dtrain, .. } = *run;
let TreeSample {
gpair,
rows,
forest_index,
} = sample;
let hist = |ghist: &GHistIndex,
backend: &dyn HistogramBackend,
reuse: Option<&ReuseSet>,
sampler: &mut ColumnSampler| {
let builder = HistTreeBuilder::new(params)
.with_rounding_seed(rounding_seed)
.with_reuse(reuse, ghist.cuts())
.with_backend(backend);
if capture_rows {
builder.build_with_leaf_rows(ghist, gpair, rows, sampler)
} else {
(builder.build(ghist, gpair, rows, sampler), Vec::new())
}
};
let (tree, leaf_rows) = match self {
Prepared::Exact(cols) => {
let builder = ExactTreeBuilder::new(params).with_reuse(reuse.as_deref());
if capture_rows {
builder.build_with_leaf_rows(cols, dtrain, gpair, rows, sampler)
} else {
(
builder.build(cols, dtrain, gpair, rows, sampler),
Vec::new(),
)
}
}
Prepared::Hist { index, backend, .. } => {
hist(index, backend.as_ref(), reuse.as_deref(), sampler)
}
Prepared::Approx { const_hess, cached } => {
let bin = || approx_index(params, dtrain, gpair, *const_hess);
let cpu = CpuBackend;
if *const_hess {
hist(cached.get_or_init(bin), &cpu, reuse.as_deref(), sampler)
} else if let Some(shared) = forest_index {
hist(shared.get_or_init(bin), &cpu, reuse.as_deref(), sampler)
} else {
hist(&bin(), &cpu, reuse.as_deref(), sampler)
}
}
};
if let Some(reuse) = reuse {
reuse.record_tree(&tree);
}
(tree, leaf_rows)
}
pub(super) fn forest_indices(
&self,
n_out: usize,
num_parallel_tree: usize,
) -> Vec<OnceLock<GHistIndex>> {
match self {
Prepared::Approx {
const_hess: false, ..
} if num_parallel_tree > 1 => (0..n_out).map(|_| OnceLock::new()).collect(),
_ => Vec::new(),
}
}
pub(super) fn samples_per_forest(&self) -> bool {
matches!(self, Prepared::Approx { .. })
}
pub(super) fn resume_approx_cache(
&self,
run: &TrainContext,
margin0: &[f32],
gpair: &mut [GradPair],
gpair_k: &mut [GradPair],
n_out: usize,
) -> Result<()> {
let TrainContext {
params,
info,
objective,
..
} = *run;
let const_hess_approx = matches!(
self,
Prepared::Approx {
const_hess: true,
..
}
);
if !const_hess_approx || !gradient_sampling(params) {
return Ok(());
}
let mut rng = round_rng(params, 0, round_salt(params));
objective.gradient_info(margin0, info, gpair);
let g0 = gather_output(gpair, gpair_k, n_out, 0);
let sampled = gradient_based_sample(g0, 1, params.subsample, &mut rng)?;
let g0 = sampled.as_ref().map_or(g0, |s| s.gpair.as_slice());
self.fill_approx_cache(run, g0);
Ok(())
}
pub(super) fn fill_approx_cache(&self, run: &TrainContext, gpair: &[GradPair]) {
if let Prepared::Approx {
const_hess: true,
cached,
} = self
{
cached.get_or_init(|| approx_index(run.params, run.dtrain, gpair, true));
}
}
}
pub(super) fn approx_index(
params: &TrainingParams,
dtrain: &DMatrix,
gpair: &[GradPair],
const_hess: bool,
) -> GHistIndex {
let cuts =
HistCuts::from_dmatrix_hessians(dtrain, params.max_bin, |row| gpair[row].hess, !const_hess);
GHistIndex::from_dmatrix(dtrain, cuts)
}
pub(super) fn prepare_builder(
params: &TrainingParams,
dtrain: &DMatrix,
const_hess: bool,
) -> Result<Prepared> {
let method = match params.tree_method {
TreeMethod::Auto | TreeMethod::Hist => TreeMethod::Hist,
TreeMethod::Exact => TreeMethod::Exact,
TreeMethod::Approx => TreeMethod::Approx,
};
if method == TreeMethod::Exact && gradient_sampling(params) {
return Err(HessboostError::invalid_param(
"sampling_method",
"`gradient_based` sampling requires `tree_method=hist` or `approx`; \
`exact` supports only `uniform`",
));
}
if method == TreeMethod::Exact && params.grow_policy == GrowPolicy::LossGuide {
return Err(HessboostError::invalid_param(
"grow_policy",
"`lossguide` growth requires `tree_method=hist`",
));
}
if params.grow_policy == GrowPolicy::Symmetric {
check_symmetric_input(method, dtrain)?;
}
Ok(match method {
TreeMethod::Hist => {
let cuts = HistCuts::from_dmatrix(dtrain, params.max_bin);
let index = GHistIndex::from_dmatrix(dtrain, cuts);
let backend = hist_backend(params, &index)?;
let rows_route_like_trees = dtrain.weights().is_none_or(|w| !w.contains(&0.0));
Prepared::Hist {
index,
backend,
rows_route_like_trees,
}
}
TreeMethod::Approx => Prepared::Approx {
const_hess,
cached: OnceLock::new(),
},
_ => Prepared::Exact(SortedColumns::from_dmatrix(dtrain)),
})
}
fn hist_backend(params: &TrainingParams, index: &GHistIndex) -> Result<Box<dyn HistogramBackend>> {
match params.device {
Device::Cpu => {
let backend: Box<dyn HistogramBackend> = Box::new(CpuBackend);
Ok(backend)
}
Device::Metal => {
#[cfg(all(target_os = "macos", feature = "metal"))]
{
let backend: Box<dyn HistogramBackend> =
Box::new(crate::backend::metal::MetalHistBackend::new(index)?);
Ok(backend)
}
#[cfg(not(all(target_os = "macos", feature = "metal")))]
{
let _ = index;
Err(HessboostError::invalid_param(
"device",
"`metal` requires building with the `metal` feature on macOS",
))
}
}
}
}