use super::dart::{dart_new_tree_weight, finish_dart, round_gradients};
use super::margins::{MarginCaches, TreeOutput};
use super::prepare::TrainContext;
use super::round::tree_eta;
use super::row_sampling::{gradient_sampling, iteration_row_subsets, make_column_sampler};
use crate::config::{BoosterKind, Device, MultiStrategy, TrainingParams, TreeMethod};
use crate::data::ghist::GHistIndex;
use crate::error::{HessboostError, Result};
use crate::model::BoostedModel;
use crate::objective::{GradPair, Loss, SplitGradient};
use crate::rng::Rng;
use crate::training::sampling::gradient_based_sample;
use crate::training::sglb::LeafRenewal;
use crate::tree::RegTree;
use crate::tree::builder::{LeafRows, MultiTreeBuilder, VectorGradients};
use crate::tree::constraints::MonotoneConstraints;
pub(super) fn vector_leaf(params: &TrainingParams, n_outputs: usize) -> bool {
params.multi_strategy == MultiStrategy::MultiOutputTree
&& params.booster != BoosterKind::GbLinear
&& n_outputs > 1
}
pub(super) fn validate(params: &TrainingParams, n_outputs: usize) -> Result<()> {
if params.multi_strategy == MultiStrategy::MultiOutputTree
&& params.booster != BoosterKind::GbLinear
&& !matches!(params.tree_method, TreeMethod::Hist | TreeMethod::Auto)
{
return Err(HessboostError::invalid_param(
"multi_strategy",
"`multi_output_tree` requires `tree_method=hist` (or `auto`)",
));
}
if vector_leaf(params, n_outputs) && params.device != Device::Cpu {
return Err(HessboostError::invalid_param(
"device",
"`metal` does not support `multi_strategy = multi_output_tree` \
(the vector-leaf builder has its own histogram loop)",
));
}
Ok(())
}
pub(crate) fn reject_split_gradient(
objective: &dyn Loss,
round: usize,
gpair: &[GradPair],
) -> Result<()> {
if objective.split_gradient(round, gpair).is_some() {
return Err(HessboostError::invalid_param(
"objective",
"reduced split gradients require `multi_strategy=multi_output_tree` \
with more than one output",
));
}
Ok(())
}
fn split_gradient(
objective: &dyn Loss,
params: &TrainingParams,
round: usize,
gpair: &[GradPair],
n_rows: usize,
) -> Result<Option<SplitGradient>> {
let Some(split) = objective.split_gradient(round, gpair) else {
return Ok(None);
};
if split.n_targets == 0 || Some(split.gpair.len()) != n_rows.checked_mul(split.n_targets) {
return Err(HessboostError::dimension_mismatch(
"split gradient length (n_rows * split n_targets)",
n_rows.saturating_mul(split.n_targets.max(1)),
split.gpair.len(),
));
}
if MonotoneConstraints::from_params(¶ms.monotone_constraints).is_active() {
return Err(HessboostError::invalid_param(
"monotone_constraints",
"monotone constraints are not supported with reduced split gradients",
));
}
Ok(Some(split))
}
pub(super) struct VectorRound<'a> {
pub(super) run: TrainContext<'a>,
pub(super) ghist: &'a GHistIndex,
pub(super) all_rows: &'a [u32],
}
pub(super) fn boost_round(
ctx: &VectorRound,
model: &mut BoostedModel,
iteration: usize,
margins: &mut MarginCaches,
gpair: &mut [GradPair],
noisy: &mut Vec<GradPair>,
) -> Result<()> {
let params = ctx.run.params;
let n = ctx.run.dtrain.n_rows();
let n_out = model.n_outputs();
let (mut rng, dropped) = round_gradients(&ctx.run, model, iteration, &margins.train, gpair);
let split = split_gradient(ctx.run.objective, params, iteration, gpair, n)?;
let weight = dart_new_tree_weight(dropped.as_ref(), params);
let structure = ctx.run.langevin.map(|langevin| {
let searched = split.as_ref().map_or(&gpair[..], |s| &s.gpair[..]);
langevin.structure_gradients(searched, iteration, noisy)
});
let grads = IterationGradients {
gpair,
split: split.as_ref(),
structure,
n_out,
iteration,
};
let row_subsets = iteration_row_subsets(params, false, ctx.run.rows, ctx.all_rows, &mut rng);
for p in 0..params.num_parallel_tree {
let rows = row_subsets.rows(p);
let (tree, leaf_rows) = fit_tree(ctx, &grads, &mut rng, rows)?;
if dropped.is_none() {
let captured =
(rows.len() == n && !gradient_sampling(params)).then_some(leaf_rows.as_slice());
margins.add_tree(&tree, TreeOutput::Vector, captured);
}
model.push_tree_weighted(tree, weight);
}
if let Some(dropped) = &dropped {
finish_dart(model, params, dropped, margins);
}
Ok(())
}
struct IterationGradients<'a> {
gpair: &'a [GradPair],
split: Option<&'a SplitGradient>,
structure: Option<&'a [GradPair]>,
n_out: usize,
iteration: usize,
}
fn fit_tree(
ctx: &VectorRound,
grads: &IterationGradients,
rng: &mut Rng,
rows: &[u32],
) -> Result<(RegTree, Vec<LeafRows>)> {
let params = ctx.run.params;
let IterationGradients {
gpair,
split,
n_out,
..
} = *grads;
let (split_gpair, n_split) = split.map_or((gpair, n_out), |s| (&s.gpair[..], s.n_targets));
let split_gpair = grads.structure.unwrap_or(split_gpair);
let sampled = if gradient_sampling(params) {
gradient_based_sample(split_gpair, n_split, params.subsample, rng)?
} else {
None
};
let sampled_value = match (&sampled, split) {
(Some(sample), Some(_)) => Some(sample.apply(gpair, n_out)),
_ => None,
};
let (split_gpair, rows) = match &sampled {
Some(sample) => (sample.gpair.as_slice(), sample.rows.as_slice()),
None => (split_gpair, rows),
};
let value = split.map(|_| sampled_value.as_deref().unwrap_or(gpair));
let mut sampler = make_column_sampler(ctx.run.dtrain, params, rng);
let grad = VectorGradients {
split: split_gpair,
n_split,
value,
n_outputs: n_out,
};
let (mut tree, leaf_rows) =
MultiTreeBuilder::new(params).build(ctx.ghist, &grad, rows, &mut sampler);
if let Some(langevin) = ctx.run.langevin {
let at = LeafRenewal {
data: ctx.run.dtrain,
gpair,
n_out,
rows,
leaf_rows: &leaf_rows,
iteration: grads.iteration,
tree: 0,
};
langevin.renew_leaves(&mut tree, TreeOutput::Vector, &at);
}
tree.scale_leaves(tree_eta(params));
Ok((tree, leaf_rows))
}