use rayon::prelude::*;
use crate::config::{ModelShrinkMode, TrainingParams};
use crate::data::DMatrix;
use crate::error::{HessboostError, Result};
use crate::objective::GradPair;
use crate::rng::{keyed_normal, stream_key};
use crate::tree::RegTree;
use crate::tree::builder::{LeafRows, xgb_calc_weight};
use crate::tree::gain::{GradStats, RegParams};
use super::margins::TreeOutput;
const STRUCTURE_STREAM: u64 = 0x5347_4C42_5F73_7472;
const LEAF_STREAM: u64 = 0x5347_4C42_5F6C_6566;
const NOISE_CHUNK: usize = 8192;
pub(super) struct Sglb {
pub(super) langevin: Option<Langevin>,
pub(super) shrink: Option<Shrink>,
}
impl Sglb {
pub(super) fn resolve(params: &TrainingParams, n_rows: usize) -> Result<Sglb> {
if params.posterior_sampling && n_rows == 0 {
return Err(HessboostError::invalid_param(
"posterior_sampling",
"needs training rows to derive its temperature and shrink rate from",
));
}
let shrink = params.effective_model_shrink(n_rows);
if params.posterior_sampling
&& let Some((rate, _)) = shrink
&& rate * params.eta >= 1.0
{
return Err(HessboostError::invalid_param(
"posterior_sampling",
format!(
"the shrink coefficient 1 - eta / (2 * rows) must stay positive, got eta {} \
for {n_rows} rows",
params.eta
),
));
}
let langevin = params.langevin_on().then(|| {
let temperature = params.effective_diffusion_temperature(n_rows);
let reg = RegParams::from_params(params);
Langevin {
sigma: params.langevin_noise_scale(temperature),
reg,
seed: params.seed,
}
});
let shrink = shrink.map(|(rate, mode)| Shrink {
mode,
rate,
eta: params.eta,
});
Ok(Sglb { langevin, shrink })
}
}
pub(super) struct Shrink {
mode: ModelShrinkMode,
rate: f64,
eta: f64,
}
impl Shrink {
pub(super) fn factor(&self, iteration: usize) -> f64 {
if iteration == 0 {
return 1.0;
}
match self.mode {
ModelShrinkMode::Constant => 1.0 - self.rate * self.eta,
ModelShrinkMode::Decreasing => 1.0 - self.rate / iteration as f64,
}
}
}
pub(super) struct Langevin {
sigma: f64,
reg: RegParams,
seed: u64,
}
pub(super) struct LeafRenewal<'a> {
pub(super) data: &'a DMatrix,
pub(super) gpair: &'a [GradPair],
pub(super) n_out: usize,
pub(super) rows: &'a [u32],
pub(super) leaf_rows: &'a [LeafRows],
pub(super) iteration: usize,
pub(super) tree: usize,
}
impl Langevin {
pub(super) fn structure_gradients<'a>(
&self,
gpair: &[GradPair],
iteration: usize,
noisy: &'a mut Vec<GradPair>,
) -> &'a [GradPair] {
let key = stream_key(&[self.seed, STRUCTURE_STREAM, iteration as u64]);
let sigma = self.sigma;
noisy.clear();
noisy.extend_from_slice(gpair);
noisy
.par_chunks_mut(NOISE_CHUNK)
.enumerate()
.for_each(|(chunk, cells)| {
let first = chunk * NOISE_CHUNK;
for (i, cell) in cells.iter_mut().enumerate() {
let z = keyed_normal(key, (first + i) as u64);
cell.grad = (f64::from(cell.grad) + sigma * z) as f32;
}
});
noisy
}
pub(super) fn renew_leaves(&self, tree: &mut RegTree, output: TreeOutput, at: &LeafRenewal) {
let (first, width) = match output {
TreeOutput::Scalar(k) => (k, 1),
TreeOutput::Vector => (0, at.n_out),
};
let mut stats = vec![GradStats::default(); tree.num_nodes() * width];
let mut add = |node: usize, row: u32| {
let cells = &at.gpair[row as usize * at.n_out + first..][..width];
for (stat, &cell) in stats[node * width..][..width].iter_mut().zip(cells) {
stat.add(GradStats::from_pair(cell));
}
};
if at.leaf_rows.is_empty() {
let leaves: Vec<usize> = at
.rows
.par_iter()
.with_min_len(1024)
.map(|&row| tree.leaf_id_with(|f| at.data.get(row as usize, f as usize)))
.collect();
for (&row, &leaf) in at.rows.iter().zip(&leaves) {
add(leaf, row);
}
} else {
for leaf in at.leaf_rows {
for &row in &leaf.rows {
add(leaf.node, row);
}
}
}
let key = stream_key(&[self.seed, LEAF_STREAM, at.iteration as u64, at.tree as u64]);
let mut values = vec![0.0f32; width];
for node in 0..tree.num_nodes() {
if !tree.node(node).is_leaf() {
continue;
}
for (out, value) in values.iter_mut().enumerate() {
let GradStats { grad, hess } = stats[node * width + out];
if hess < self.reg.min_child_weight || hess <= 0.0 {
*value = 0.0;
continue;
}
let scale = self.sigma * (hess.abs() + self.reg.lambda).sqrt();
let z = keyed_normal(key, (node * width + out) as u64);
*value = xgb_calc_weight(GradStats::new(grad + scale * z, hess), &self.reg) as f32;
}
match output {
TreeOutput::Scalar(_) => tree.set_leaf_value(node, values[0]),
TreeOutput::Vector => tree.set_leaf_vector(node, &values),
}
}
}
}