use crate::config::{Refresh, TrainingParams};
use crate::data::DMatrix;
use crate::objective::GradPair;
use crate::tree::RegTree;
use crate::tree::gain::{GradStats, RegParams, calc_gain, calc_weight};
use rayon::prelude::*;
const REFRESH_BLOCK_ROWS: usize = 4096;
const REFRESH_BATCH_PER_THREAD: usize = 4;
pub(super) fn refresh_tree(
tree: &mut RegTree,
data: &DMatrix,
gpair: &[GradPair],
params: &TrainingParams,
refresh: Refresh,
learning_rate: f32,
) {
let reg = &RegParams {
min_child_weight: 0.0,
..RegParams::from_params(params)
};
let stats = node_stats(tree, data, gpair);
for nid in 0..tree.num_nodes() {
let node = *tree.node(nid);
tree.set_sum_hess(nid, stats[nid].hess as f32);
if node.is_leaf() {
if refresh.refresh_leaf() {
let base_weight = calc_weight(stats[nid], reg) as f32;
tree.set_leaf_value(nid, base_weight * learning_rate);
}
} else {
let gain = calc_gain(stats[node.left as usize], reg)
+ calc_gain(stats[node.right as usize], reg)
- calc_gain(stats[nid], reg);
tree.set_split_gain(nid, gain as f32);
}
}
}
fn node_stats(tree: &RegTree, data: &DMatrix, gpair: &[GradPair]) -> Vec<GradStats> {
let batch_blocks = rayon::current_num_threads().max(1) * REFRESH_BATCH_PER_THREAD;
node_stats_batched(tree, data, gpair, batch_blocks)
}
fn node_stats_batched(
tree: &RegTree,
data: &DMatrix,
gpair: &[GradPair],
batch_blocks: usize,
) -> Vec<GradStats> {
let n_nodes = tree.num_nodes();
let block_stats = |rows: std::ops::Range<usize>| {
let mut stats = vec![GradStats::default(); n_nodes];
for row in rows {
let gp = GradStats::from_pair(gpair[row]);
let mut nid = 0;
stats[nid].add(gp);
while !tree.node(nid).is_leaf() {
let feature = tree.node(nid).split_feature as usize;
nid = tree.child(nid, data.get(row, feature));
stats[nid].add(gp);
}
}
stats
};
let n = data.n_rows();
let n_blocks = n.div_ceil(REFRESH_BLOCK_ROWS);
let batch = batch_blocks.max(1);
let mut total = vec![GradStats::default(); n_nodes];
for first in (0..n_blocks).step_by(batch) {
let blocks: Vec<Vec<GradStats>> = (first..(first + batch).min(n_blocks))
.into_par_iter()
.map(|b| block_stats(b * REFRESH_BLOCK_ROWS..((b + 1) * REFRESH_BLOCK_ROWS).min(n)))
.collect();
for block in blocks {
for (acc, s) in total.iter_mut().zip(block) {
acc.add(s);
}
}
}
total
}
#[cfg(test)]
mod tests {
use super::*;
use crate::tree::{ChildLeaf, SplitRule};
fn stump() -> RegTree {
let mut t = RegTree::with_root(99.0);
t.expand(
0,
SplitRule::numeric(0, 0.5, true),
ChildLeaf::new(7.0, 99.0),
ChildLeaf::new(-7.0, 99.0),
);
t
}
#[test]
fn refresh_recomputes_cover_gain_and_leaves_from_the_routed_rows() {
let data = DMatrix::from_dense(&[0.1, 0.2, 0.9, f32::NAN], 4, 1).unwrap();
let gpair = [
GradPair::new(1.0, 1.0),
GradPair::new(2.0, 1.0),
GradPair::new(-3.0, 2.0),
GradPair::new(0.5, 0.5),
];
let params = TrainingParams::builder()
.lambda(1.0)
.min_child_weight(3.0)
.build()
.unwrap();
let reg = RegParams {
min_child_weight: 0.0,
..RegParams::from_params(¶ms)
};
let mut tree = stump();
refresh_tree(&mut tree, &data, &gpair, ¶ms, Refresh::default(), 0.25);
let (left, right, root) = (
GradStats::new(3.5, 2.5),
GradStats::new(-3.0, 2.0),
GradStats::new(0.5, 4.5),
);
assert_eq!(tree.node(0).sum_hess, 4.5);
assert_eq!(tree.node(1).sum_hess, 2.5);
assert_eq!(tree.node(2).sum_hess, 2.0);
let gain = calc_gain(left, ®) + calc_gain(right, ®) - calc_gain(root, ®);
assert_eq!(tree.node(0).split_gain, gain as f32);
assert_eq!(tree.node(1).leaf_value, (-3.5f64 / 3.5) as f32 * 0.25);
assert_eq!(tree.node(2).leaf_value, (3.0f64 / 3.0) as f32 * 0.25);
let mut kept = stump();
refresh_tree(
&mut kept,
&data,
&gpair,
¶ms,
Refresh::stats_only(),
0.25,
);
assert_eq!(kept.node(1).leaf_value, 7.0);
assert_eq!(kept.node(2).leaf_value, -7.0);
assert_eq!(kept.node(0).sum_hess, 4.5);
assert_eq!(kept.node(0).split_gain, gain as f32);
}
#[test]
fn a_leaf_no_row_reaches_gets_zero_weight() {
let data = DMatrix::from_dense(&[0.1, 0.2], 2, 1).unwrap();
let gpair = [GradPair::new(1.0, 1.0), GradPair::new(1.0, 1.0)];
let mut tree = stump();
refresh_tree(
&mut tree,
&data,
&gpair,
&TrainingParams::default(),
Refresh::default(),
0.3,
);
assert_eq!(tree.node(2).sum_hess, 0.0);
assert_eq!(tree.node(2).leaf_value, 0.0);
}
#[test]
fn batched_block_reduction_matches_for_every_batch_size() {
let n = 5 * REFRESH_BLOCK_ROWS + 123;
let x: Vec<f32> = (0..n).map(|i| ((i * 37) % 101) as f32 / 101.0).collect();
let data = DMatrix::from_dense(&x, n, 1).unwrap();
let gpair: Vec<GradPair> = (0..n)
.map(|i| {
let g = ((i * 7919) % 1000) as f32 * 1e-3 - 0.37;
GradPair::new(g, 0.1 + (i % 13) as f32)
})
.collect();
let tree = stump();
let unbatched = node_stats_batched(&tree, &data, &gpair, usize::MAX);
assert!(unbatched.iter().all(|s| s.hess > 0.0));
for batch in [0, 1, 2, 4, 6] {
let stats = node_stats_batched(&tree, &data, &gpair, batch);
for (a, b) in stats.iter().zip(&unbatched) {
assert_eq!(
(a.grad.to_bits(), a.hess.to_bits()),
(b.grad.to_bits(), b.hess.to_bits()),
"batch {batch}"
);
}
}
}
}