use super::fast::fast_pairs;
use super::{EBM_BOULEVARD_SALT, Grown, Hook, StageFit, Term, gradients, grow, parallel};
use crate::error::Result;
use crate::training::boulevard::{Recursion, RoundRequest, Schedule, tree_rows};
use crate::training::prepare::{Prepared, TrainContext};
use crate::training::row_sampling::sample_rows;
use crate::tree::RegTree;
use rayon::prelude::*;
pub(super) fn boulevard(
run: &TrainContext,
prepared: &Prepared,
mains: &[Vec<u32>],
mu: f64,
rounds: usize,
hook: &mut Hook,
) -> Result<Grown> {
let params = run.params;
let n = run.dtrain.n_rows();
let main_terms: Vec<Term> = mains
.iter()
.enumerate()
.map(|(t, f)| (t as u32, f.as_slice()))
.collect();
let base = vec![mu; n];
let StageFit { mut trees, fitted } =
boulevard_stage(run, prepared, &main_terms, &base, (0, rounds), hook)?;
let mut pairs = Vec::new();
if params.ebm_settings().interactions() > 0 && !hook.stopped {
let margins: Vec<f32> = fitted.iter().map(|&m| m as f32).collect();
pairs = fast_pairs(
run,
&gradients(run, &margins, rounds),
params.ebm_settings().interactions(),
);
let pair_terms: Vec<Term> = pairs
.iter()
.enumerate()
.map(|(k, f)| ((mains.len() + k) as u32, f.as_slice()))
.collect();
let stage = boulevard_stage(run, prepared, &pair_terms, &fitted, (1, rounds), hook)?;
trees.extend(stage.trees);
}
Ok(Grown { trees, pairs })
}
fn boulevard_stage(
run: &TrainContext,
prepared: &Prepared,
terms: &[Term],
base: &[f64],
stage: (u64, usize),
hook: &mut Hook,
) -> Result<StageFit> {
let (stage, rounds) = stage;
let TrainContext { params, dtrain, .. } = *run;
let n = dtrain.n_rows();
let schedule = Schedule {
dropout: 0.0,
learning_rate: params.eta,
truncation: None,
parallel: 1,
seed: params.seed,
salt: EBM_BOULEVARD_SALT ^ stage,
};
let mut recursion = Recursion::new(schedule, n);
let mut trees: Vec<(u32, RegTree)> = Vec::with_capacity(rounds * terms.len());
let mut total = vec![0.0f64; n];
for round in 0..rounds {
let iteration = stage as usize * rounds + round;
recursion.step(|request| {
let RoundRequest { offsets, rng, .. } = request;
let margins: Vec<f32> = base
.iter()
.zip(&offsets[0])
.map(|(&b, &o)| (b + o) as f32)
.collect();
let gpair = gradients(run, &margins, iteration);
let rows: Vec<Vec<u32>> = terms
.iter()
.map(|_| sample_rows(n, params, run.rows, rng))
.collect();
let seeds: Vec<u64> = terms.iter().map(|_| rng.next_u64()).collect();
prepared.fill_approx_cache(run, &gpair);
let build = |k: usize| grow(run, prepared, &gpair, &rows[k], terms[k].1, seeds[k]);
let grown: Vec<RegTree> = if parallel(params) && terms.len() > 1 {
(0..terms.len()).into_par_iter().map(build).collect()
} else {
(0..terms.len()).map(build).collect()
};
let mut round_sum = vec![0.0f64; n];
for (&(term, _), mut tree) in terms.iter().zip(grown) {
let preds = tree_rows(&tree, dtrain);
let mean = preds.iter().map(|&v| f64::from(v)).sum::<f64>() / n as f64;
tree.shift_leaves(-mean as f32);
for (s, v) in round_sum.iter_mut().zip(tree_rows(&tree, dtrain)) {
*s += f64::from(v);
}
trees.push((term, tree));
}
for (t, &s) in total.iter_mut().zip(&round_sum) {
*t += s;
}
Ok(vec![round_sum.iter().map(|&s| s as f32).collect()])
})?;
if !hook.next() {
break;
}
}
let scale = recursion.scale();
for (_, tree) in &mut trees {
tree.scale_leaves(scale as f32);
}
let fitted = base
.iter()
.zip(&total)
.map(|(&b, &t)| b + scale * t)
.collect();
Ok(StageFit { trees, fitted })
}