use super::dart::{dart_new_tree_weight, finish_dart, round_gradients};
use super::margins::{MarginCaches, TreeOutput};
use super::prepare::{Prepared, TrainContext, TreeSample, approx_index};
use super::row_sampling::{gradient_sampling, iteration_row_subsets, make_column_sampler};
use crate::config::{Device, Refresh, TrainingParams};
use crate::data::ghist::GHistIndex;
use crate::error::Result;
use crate::model::BoostedModel;
use crate::objective::GradPair;
use crate::rng::Rng;
use crate::training::multi_output;
use crate::training::refresh::refresh_tree;
use crate::training::sampling::{GradientSample, gradient_based_sample};
use crate::training::sglb::LeafRenewal;
use crate::tree::RegTree;
use crate::tree::builder::LeafRows;
use crate::tree::reuse::ReuseSet;
use crate::tree::sampler::ColumnSampler;
use rayon::prelude::*;
use std::sync::OnceLock;
pub(super) struct RoundState<'a> {
pub(super) model: BoostedModel,
pub(super) margins: MarginCaches<'a>,
pub(super) gpair: Vec<GradPair>,
pub(super) gpair_k: Vec<GradPair>,
pub(super) reuse: Option<ReuseSet>,
pub(super) noisy_gpair: Vec<GradPair>,
pub(super) all_rows: Vec<u32>,
}
pub(super) fn refresh_round(
run: &TrainContext,
queue: &mut [RegTree],
refresh: Refresh,
iteration: usize,
state: &mut RoundState,
) -> Result<()> {
let TrainContext {
params,
dtrain,
info,
objective,
..
} = *run;
let n_out = objective.n_outputs();
let parallel = params.num_parallel_tree;
objective.gradient_info_at(&state.margins.train, info, &mut state.gpair, iteration);
multi_output::reject_split_gradient(objective, iteration, &state.gpair)?;
let per_iteration = n_out * parallel;
for slot in 0..per_iteration {
let k = slot / parallel;
let gk = gather_output(&state.gpair, &mut state.gpair_k, n_out, k);
let mut tree = std::mem::replace(
&mut queue[iteration * per_iteration + slot],
RegTree::with_root(0.0),
);
refresh_tree(&mut tree, dtrain, gk, params, refresh, tree_eta(params));
state.margins.add_tree(&tree, TreeOutput::Scalar(k), None);
state.model.push_tree_weighted(tree, 1.0);
}
Ok(())
}
pub(super) fn grow_round(
run: &TrainContext,
prepared: &Prepared,
iteration: usize,
state: &mut RoundState,
) -> Result<()> {
let TrainContext {
params,
dtrain,
objective,
..
} = *run;
let n = dtrain.n_rows();
let n_out = objective.n_outputs();
let parallel = params.num_parallel_tree;
let (mut rng, dropped) = round_gradients(
run,
&state.model,
iteration,
&state.margins.train,
&mut state.gpair,
);
multi_output::reject_split_gradient(objective, iteration, &state.gpair)?;
let weight = dart_new_tree_weight(dropped.as_ref(), params);
let structure: &[GradPair] = match run.langevin {
Some(langevin) => {
langevin.structure_gradients(&state.gpair, iteration, &mut state.noisy_gpair)
}
None => &state.gpair,
};
let row_subsets = iteration_row_subsets(
params,
prepared.samples_per_forest(),
run.rows,
&state.all_rows,
&mut rng,
);
let mut forest_sample = None;
let forest_indices = prepared.forest_indices(n_out, parallel);
let grow = GrowRound {
run,
prepared,
gpair: structure,
clean_gpair: &state.gpair,
n_out,
iteration,
forest_indices: &forest_indices,
};
let slots: Vec<TreeSlot> = (0..n_out * parallel)
.map(|slot| {
let row_subset = row_subsets.rows(slot % parallel);
let (routed, linear_rows) = match prepared {
Prepared::Hist {
rows_route_like_trees,
..
} => (true, params.linear_tree.is_some() && *rows_route_like_trees),
Prepared::Exact(_) => (true, false),
Prepared::Approx { .. } => (false, false),
};
let margin_rows = routed
&& dropped.is_none()
&& row_subset.len() == n
&& !gradient_sampling(params)
&& (params.linear_tree.is_none() || linear_rows);
TreeSlot {
output: slot / parallel,
parallel: slot % parallel,
rows: row_subset,
capture_rows: margin_rows || linear_rows || (routed && run.langevin.is_some()),
margin_rows,
}
})
.collect();
let trees: Vec<(RegTree, Vec<LeafRows>)> = if slots.len() > 1
&& state.reuse.is_none()
&& !gradient_sampling(params)
&& params.device == Device::Cpu
&& rayon::current_num_threads() > 1
{
prepared.fill_approx_cache(run, gather_output(structure, &mut state.gpair_k, n_out, 0));
let draws: Vec<(ColumnSampler, u64)> = slots
.iter()
.map(|_| {
let sampler = make_column_sampler(dtrain, params, &mut rng);
(sampler, quantization_seed(params, &mut rng))
})
.collect();
let gathered = gather_outputs(structure, n_out);
let output_gpair = |k: usize| {
if n_out == 1 {
structure
} else {
&gathered[k * n..(k + 1) * n]
}
};
for (k, index) in forest_indices.iter().enumerate() {
index.get_or_init(|| approx_index(params, dtrain, output_gpair(k), false));
}
slots
.par_iter()
.zip(draws)
.map(|(slot, (mut sampler, rounding_seed))| {
let gk = output_gpair(slot.output);
let sample = TreeSample {
gpair: gk,
rows: slot.rows,
forest_index: grow.forest_index(slot.output),
};
grow_sampled_tree(&grow, slot, sample, &mut sampler, rounding_seed, None)
})
.collect()
} else {
slots
.iter()
.map(|slot| {
fit_output_tree(
&grow,
slot,
&mut state.gpair_k,
&mut rng,
&mut forest_sample,
state.reuse.as_mut(),
)
})
.collect::<Result<_>>()?
};
for (slot, (tree, leaf_rows)) in slots.iter().zip(trees) {
if dropped.is_none() {
let captured = slot.margin_rows.then_some(leaf_rows.as_slice());
state
.margins
.add_tree(&tree, TreeOutput::Scalar(slot.output), captured);
}
state.model.push_tree_weighted(tree, weight);
}
if let Some(dropped) = &dropped {
finish_dart(&mut state.model, params, dropped, &mut state.margins);
}
Ok(())
}
pub(super) enum RoundPlan {
Grow(Prepared),
Refresh(Vec<RegTree>, Refresh),
}
pub(super) fn tree_eta(params: &TrainingParams) -> f32 {
params.eta as f32 / params.num_parallel_tree as f32
}
pub(super) fn gather_output<'a>(
gpair: &'a [GradPair],
scratch: &'a mut [GradPair],
n_out: usize,
k: usize,
) -> &'a [GradPair] {
if n_out == 1 {
gpair
} else {
for (r, dst) in scratch.iter_mut().enumerate() {
*dst = gpair[r * n_out + k];
}
scratch
}
}
fn gather_outputs(gpair: &[GradPair], n_out: usize) -> Vec<GradPair> {
if n_out == 1 {
return Vec::new();
}
let n = gpair.len() / n_out;
let mut out = vec![GradPair::default(); gpair.len()];
out.par_chunks_exact_mut(n)
.enumerate()
.for_each(|(k, column)| {
for (r, dst) in column.iter_mut().enumerate() {
*dst = gpair[r * n_out + k];
}
});
out
}
struct GrowRound<'a> {
run: &'a TrainContext<'a>,
prepared: &'a Prepared,
gpair: &'a [GradPair],
clean_gpair: &'a [GradPair],
n_out: usize,
iteration: usize,
forest_indices: &'a [OnceLock<GHistIndex>],
}
impl GrowRound<'_> {
fn forest_index(&self, output: usize) -> Option<&OnceLock<GHistIndex>> {
self.forest_indices.get(output)
}
}
struct TreeSlot<'a> {
output: usize,
parallel: usize,
rows: &'a [u32],
capture_rows: bool,
margin_rows: bool,
}
fn fit_output_tree(
grow: &GrowRound,
slot: &TreeSlot,
scratch: &mut [GradPair],
rng: &mut Rng,
forest_sample: &mut Option<GradientSample>,
reuse: Option<&mut ReuseSet>,
) -> Result<(RegTree, Vec<LeafRows>)> {
let TrainContext { params, dtrain, .. } = *grow.run;
let (prepared, n_out) = (grow.prepared, grow.n_out);
let gk: &[GradPair] = gather_output(grow.gpair, scratch, n_out, slot.output);
let own;
let sampled = if !gradient_sampling(params) {
None
} else if prepared.samples_per_forest() {
if slot.parallel == 0 {
*forest_sample = gradient_based_sample(gk, 1, params.subsample, rng)?;
}
forest_sample.as_ref()
} else {
own = gradient_based_sample(gk, 1, params.subsample, rng)?;
own.as_ref()
};
let (gk, rows) = match sampled {
Some(s) => (s.gpair.as_slice(), s.rows.as_slice()),
None => (gk, slot.rows),
};
let mut sampler = make_column_sampler(dtrain, params, rng);
let rounding_seed = quantization_seed(params, rng);
let sample = TreeSample {
gpair: gk,
rows,
forest_index: grow.forest_index(slot.output),
};
Ok(grow_sampled_tree(
grow,
slot,
sample,
&mut sampler,
rounding_seed,
reuse,
))
}
fn grow_sampled_tree(
grow: &GrowRound,
slot: &TreeSlot,
sample: TreeSample,
sampler: &mut ColumnSampler,
rounding_seed: u64,
reuse: Option<&mut ReuseSet>,
) -> (RegTree, Vec<LeafRows>) {
let TrainContext { params, dtrain, .. } = *grow.run;
let TreeSample {
gpair: gk, rows, ..
} = sample;
let (mut tree, leaf_rows) = grow.prepared.build_tree(
grow.run,
sample,
sampler,
reuse,
rounding_seed,
slot.capture_rows,
);
if let Some(langevin) = grow.run.langevin {
let at = LeafRenewal {
data: dtrain,
gpair: grow.clean_gpair,
n_out: grow.n_out,
rows,
leaf_rows: &leaf_rows,
iteration: grow.iteration,
tree: slot.output * params.num_parallel_tree + slot.parallel,
};
langevin.renew_leaves(&mut tree, TreeOutput::Scalar(slot.output), &at);
}
if let Some(linear_tree) = params.linear_tree
&& grow.iteration > 0
{
let lambda = linear_tree.lambda();
if leaf_rows.is_empty() {
crate::tree::linear_fit::fit_linear_leaves(&mut tree, dtrain, gk, rows, lambda);
} else {
crate::tree::linear_fit::fit_captured_linear_leaves(
&mut tree, dtrain, gk, &leaf_rows, lambda,
);
}
}
tree.scale_leaves(tree_eta(params));
(tree, leaf_rows)
}
fn quantization_seed(params: &TrainingParams, rng: &mut Rng) -> u64 {
if params.quantized.is_some() {
rng.next_u64()
} else {
0
}
}