use super::check_data;
use crate::data::DMatrix;
use crate::ebm::{EbmBoulevard, EbmInfo};
use crate::error::{HessboostError, Result};
use crate::model::BoostedModel;
use crate::rng::Rng;
use crate::training::boulevard::{Recursion, RoundRequest, Schedule};
use crate::tree::RegTree;
const REFIT_SALT: u64 = 0x1_0E57;
fn node_counts(tree: &RegTree, leaf_counts: &[usize]) -> Vec<usize> {
let nodes = tree.nodes();
let mut counts = leaf_counts.to_vec();
let mut stack = vec![(0usize, false)];
while let Some((id, expanded)) = stack.pop() {
let node = &nodes[id];
if node.is_leaf() {
continue;
}
let (l, r) = (node.left as usize, node.right as usize);
if expanded {
counts[id] = counts[l] + counts[r];
} else {
stack.push((id, true));
stack.push((l, false));
stack.push((r, false));
}
}
counts
}
struct LeafIds {
ids: Vec<u32>,
rows: usize,
trees: usize,
}
impl LeafIds {
fn new(model: &BoostedModel, values: &DMatrix) -> Result<Self> {
Ok(LeafIds {
ids: model.predict_leaf(values, ..)?.into_vec(),
rows: values.n_rows(),
trees: model.num_trees(),
})
}
fn leaf(&self, row: usize, tree: usize) -> usize {
self.ids[row * self.trees + tree] as usize
}
}
struct LeafSums {
sums: Vec<f64>,
counts: Vec<usize>,
}
impl LeafSums {
fn accumulate(
leaves: &LeafIds,
tree: usize,
nodes: usize,
in_bag: &[bool],
residual: impl Fn(usize) -> f64,
) -> Self {
let mut sums = vec![0.0f64; nodes];
let mut counts = vec![0usize; nodes];
for (row, &keep) in in_bag.iter().enumerate() {
if keep {
let leaf = leaves.leaf(row, tree);
sums[leaf] += residual(row);
counts[leaf] += 1;
}
}
LeafSums { sums, counts }
}
fn leaf_values(&self, tree: &RegTree, reg_lambda: f64) -> Vec<f64> {
let mut values = vec![0.0f64; tree.num_nodes()];
for (id, node) in tree.nodes().iter().enumerate() {
let denom = self.counts[id] as f64 + reg_lambda;
if node.is_leaf() && denom > 0.0 {
values[id] = self.sums[id] / denom;
}
}
values
}
}
fn update_covers(trees: &mut [RegTree], leaves: &LeafIds) {
for (t, tree) in trees.iter_mut().enumerate() {
let mut leaf_counts = vec![0usize; tree.num_nodes()];
for row in 0..leaves.rows {
leaf_counts[leaves.leaf(row, t)] += 1;
}
for (id, &c) in node_counts(tree, &leaf_counts).iter().enumerate() {
tree.set_sum_hess(id, c as f32);
}
}
}
pub fn honest_refit(model: &BoostedModel, values: &DMatrix) -> Result<BoostedModel> {
if let Some(info) = model.ebm()
&& let Some(settings) = info.boulevard
{
return ebm_refit(model, info, settings, values);
}
let info = *model.boulevard().ok_or_else(|| {
HessboostError::incompatible_model(
"model",
"not a Boulevard fit: train it with `booster = boulevard` (or `booster = ebm` with \
`ebm_boulevard`)",
)
})?;
let labels = finite_labels(model, values)?;
let n = values.n_rows();
let mu = if info.intercept_from_labels {
(labels.iter().map(|&y| f64::from(y)).sum::<f64>() / n as f64) as f32
} else {
model.base_score()
};
let leaves = LeafIds::new(model, values)?;
let parallel = model.num_parallel_tree();
let schedule = Schedule::from_info(&info, parallel, REFIT_SALT);
let mut refit = model.clone();
let mut recursion = Recursion::new(schedule, n);
let trees = refit.trees_mut();
for _ in 0..model.num_boost_rounds() {
recursion.step(|request| {
let RoundRequest {
index,
first_slot,
offsets,
rng,
} = request;
let mut preds = Vec::with_capacity(offsets.len());
for (j, offset) in offsets.iter().enumerate() {
let t = index * parallel + first_slot + j;
let in_bag = row_sample(n, info.subsample, rng);
let tree = &mut trees[t];
let sums = LeafSums::accumulate(&leaves, t, tree.num_nodes(), &in_bag, |row| {
f64::from(labels[row]) - f64::from(mu) - offset[row]
});
let values: Vec<f32> = sums
.leaf_values(tree, info.reg_lambda)
.into_iter()
.map(|v| v as f32)
.collect();
for (id, &v) in values.iter().enumerate() {
if tree.nodes()[id].is_leaf() {
tree.set_leaf_value(id, v);
}
}
preds.push((0..n).map(|row| values[leaves.leaf(row, t)]).collect());
}
Ok(preds)
})?;
}
let scale = recursion.scale() as f32;
for tree in trees.iter_mut() {
tree.scale_leaves(scale);
}
update_covers(trees, &leaves);
refit.set_base_scores(vec![mu]);
refit.set_best_iteration(None);
if refit
.trees()
.iter()
.any(|t| t.nodes().iter().any(|n| !n.leaf_value.is_finite()))
{
return Err(HessboostError::invalid_data(
"values",
"the refitted leaves overflow f32",
));
}
Ok(refit)
}
fn finite_labels<'v>(model: &BoostedModel, values: &'v DMatrix) -> Result<&'v [f32]> {
check_data(model, values, "values", true)?;
let labels = values.labels().unwrap_or_default();
if labels.iter().any(|y| !y.is_finite()) {
return Err(HessboostError::invalid_data(
"values",
"labels must be finite",
));
}
Ok(labels)
}
fn ebm_refit(
model: &BoostedModel,
info: &EbmInfo,
settings: EbmBoulevard,
values: &DMatrix,
) -> Result<BoostedModel> {
let labels = finite_labels(model, values)?;
let n = values.n_rows();
let mu = labels.iter().map(|&y| f64::from(y)).sum::<f64>() / n as f64;
let leaves = LeafIds::new(model, values)?;
let mut refit = model.clone();
let mut base = vec![mu; n];
for stage in info.stages()? {
let terms = stage.terms.len();
let schedule = Schedule {
dropout: 0.0,
learning_rate: settings.learning_rate,
truncation: None,
parallel: 1,
seed: 0,
salt: REFIT_SALT ^ stage.index as u64,
};
let mut recursion = Recursion::new(schedule, n);
let mut total = vec![0.0f64; n];
let trees = refit.trees_mut();
for round in 0..stage.rounds {
recursion.step(|request| {
let RoundRequest { offsets, rng, .. } = request;
let mut round_sum = vec![0.0f64; n];
for k in 0..terms {
let t = stage.trees.start + round * terms + k;
let in_bag = row_sample(n, settings.subsample, rng);
let tree = &mut trees[t];
let sums = LeafSums::accumulate(&leaves, t, tree.num_nodes(), &in_bag, |row| {
f64::from(labels[row]) - base[row] - offsets[0][row]
});
let leaf_values = sums.leaf_values(tree, settings.reg_lambda);
let mean = (0..n)
.map(|row| leaf_values[leaves.leaf(row, t)])
.sum::<f64>()
/ n as f64;
for (id, v) in leaf_values.iter().enumerate() {
if tree.nodes()[id].is_leaf() {
tree.set_leaf_value(id, (v - mean) as f32);
}
}
for (row, s) in round_sum.iter_mut().enumerate() {
*s += f64::from(tree.nodes()[leaves.leaf(row, t)].leaf_value);
}
}
for (t, &s) in total.iter_mut().zip(&round_sum) {
*t += s;
}
Ok(vec![round_sum.iter().map(|&s| s as f32).collect()])
})?;
}
let scale = recursion.scale();
for tree in &mut trees[stage.trees] {
tree.scale_leaves(scale as f32);
}
for (b, &t) in base.iter_mut().zip(&total) {
*b += scale * t;
}
}
update_covers(refit.trees_mut(), &leaves);
refit.set_base_scores(vec![mu as f32]);
if refit
.trees()
.iter()
.any(|t| t.nodes().iter().any(|n| !n.leaf_value.is_finite()))
{
return Err(HessboostError::invalid_data(
"values",
"the refitted leaves overflow f32",
));
}
let mut refit_info = info.clone();
refit_info.term_means = refit_info.term_means_on(&refit, values);
refit.set_ebm(Some(refit_info));
Ok(refit)
}
fn row_sample(n: usize, ratio: f64, rng: &mut Rng) -> Vec<bool> {
if ratio >= 1.0 {
return vec![true; n];
}
let mut keep: Vec<bool> = (0..n).map(|_| rng.f64() < ratio).collect();
if !keep.iter().any(|&k| k) {
keep[rng.range(0..n)] = true;
}
keep
}