use crate::config::TrainingParams;
use crate::data::DMatrix;
use crate::data::ghist::{Bins, GHistIndex};
use crate::data::quantile::HistCuts;
use crate::error::{HessboostError, Result};
use crate::model::BoostedModel;
use crate::objective::{GradPair, Loss};
use crate::training::margins::{TreeOutput, add_tree_margins};
use crate::tree::RegTree;
use crate::tree::gain::{GradStats, RegParams, calc_weight};
#[derive(Debug, Clone)]
pub(super) struct Cache {
pub(super) tolerance: f64,
pub(super) cuts: HistCuts,
pub(super) dense: bool,
pub(super) trees: Vec<TreeCache>,
}
#[derive(Debug, Clone)]
pub(super) struct TreeCache {
pub(super) nodes: Vec<NodeCache>,
pub(super) grads: Vec<GradPair>,
}
#[derive(Debug, Clone, Default)]
pub(super) struct NodeCache {
pub(super) stats: GradStats,
pub(super) hist: Vec<GradStats>,
}
impl Cache {
pub(super) fn build(
model: &BoostedModel,
params: &TrainingParams,
data: &DMatrix,
tolerance: f64,
) -> Result<Self> {
let cuts = HistCuts::from_dmatrix(data, params.max_bin);
let ghist = GHistIndex::from_dmatrix(data, cuts.clone());
let objective = params.loss(1)?;
let mut margins = vec![model.base_scores()[0]; data.n_rows()];
let mut trees = Vec::with_capacity(model.num_trees());
for (t, tree) in model.trees().iter().enumerate() {
let grads = gradients(objective.as_ref(), data, &margins);
let mut nodes = vec![NodeCache::default(); tree.num_nodes()];
accumulate(
tree,
0,
data,
&ghist,
grads.iter().copied().enumerate(),
&mut nodes,
);
if t == 0 {
check_newton_leaves(tree, &nodes, params)?;
}
add_tree_margins(tree, data, &mut margins, 1, TreeOutput::Scalar(0));
trees.push(TreeCache { nodes, grads });
}
let dense = ghist.dense_stride().is_some();
Ok(Cache {
tolerance,
cuts,
dense,
trees,
})
}
}
fn check_newton_leaves(tree: &RegTree, nodes: &[NodeCache], params: &TrainingParams) -> Result<()> {
let reg = RegParams::from_params(params);
let eta = params.eta as f32;
for (nid, cache) in nodes.iter().enumerate() {
let node = tree.node(nid);
if !node.is_leaf() || cache.stats.hess <= 0.0 {
continue;
}
let expected = (calc_weight(cache.stats, ®) as f32) * eta;
let stored = node.leaf_value;
let close = (stored - expected).abs() <= 1e-3 * (1.0 + expected.abs());
if !close {
return Err(HessboostError::incompatible_model(
"model",
format!(
"leaf {nid} of the first tree holds {stored}, but these parameters give \
{expected} on this data: the approximate mode updates only models \
trained with these parameters on this data (not, for example, an \
imported LightGBM model, whose first tree carries the label average)"
),
));
}
}
Ok(())
}
fn gradients(objective: &dyn Loss, data: &DMatrix, margins: &[f32]) -> Vec<GradPair> {
let mut out = vec![GradPair::default(); data.n_rows()];
objective.gradient_info(margins, &data.info(), &mut out);
out
}
pub(super) fn stats_of(g: GradPair) -> GradStats {
GradStats::new(f64::from(g.grad), f64::from(g.hess))
}
fn add_index_bins(hist: &mut [GradStats], ghist: &GHistIndex, row: usize, g: GradStats) {
let (s, e) = (ghist.row_ptr()[row], ghist.row_ptr()[row + 1]);
match ghist.bins() {
Bins::U16(b) => b[s..e].iter().for_each(|&bin| hist[bin as usize].add(g)),
Bins::U32(b) => b[s..e].iter().for_each(|&bin| hist[bin as usize].add(g)),
}
}
pub(super) fn accumulate(
tree: &RegTree,
root: usize,
data: &DMatrix,
ghist: &GHistIndex,
rows: impl Iterator<Item = (usize, GradPair)>,
out: &mut [NodeCache],
) {
let bins = ghist.total_bins();
for (row, g) in rows {
let g = stats_of(g);
let mut nid = root;
loop {
let node = tree.node(nid);
out[nid].stats.add(g);
if node.is_leaf() {
break;
}
if out[nid].hist.is_empty() {
out[nid].hist = vec![GradStats::default(); bins];
}
add_index_bins(&mut out[nid].hist, ghist, row, g);
nid = tree.child(nid, data.get(row, node.split_feature as usize));
}
}
}
pub(super) fn clear_subtree(tree: &RegTree, root: usize, nodes: &mut [NodeCache]) {
let mut stack = vec![root];
while let Some(nid) = stack.pop() {
nodes[nid] = NodeCache::default();
let node = tree.node(nid);
if !node.is_leaf() {
stack.push(node.left as usize);
stack.push(node.right as usize);
}
}
}