use rayon::prelude::*;
use super::kernel::Kernel;
use crate::data::DMatrix;
use crate::ebm::grid::TermGrid;
use crate::error::{HessboostError, Result};
use crate::tree::RegTree;
pub(super) struct TermPart<'a> {
pub(super) grid: TermGrid,
trees: Vec<&'a RegTree>,
cell_of_row: Vec<u32>,
weight: Vec<f64>,
col_mean: Vec<f64>,
}
impl<'a> TermPart<'a> {
pub(super) fn new(
trees: Vec<&'a RegTree>,
features: &[u32],
train: &DMatrix,
kappa: f64,
) -> Result<Self> {
let grid = TermGrid::new(&trees, features);
let n = train.n_rows();
let cell_of_row: Vec<u32> = (0..n)
.into_par_iter()
.with_min_len(1024)
.map(|row| grid.cell_of_row(train, row) as u32)
.collect();
let mut counts = vec![0.0; grid.len()];
for &c in &cell_of_row {
counts[c as usize] += 1.0;
}
let prefix = grid.prefix(&counts);
let n_trees = trees.len() as f64;
let mut weight = Vec::with_capacity(grid.leaves.len());
for (t, tree) in trees.iter().enumerate() {
for leaf in &grid.leaves[grid.leaf_start[t]..grid.leaf_start[t + 1]] {
let rows: f64 = leaf.boxes.iter().map(|&b| grid.box_sum(&prefix, b)).sum();
let cover = f64::from(tree.node(leaf.node as usize).sum_hess);
if rows < cover {
return Err(HessboostError::invalid_data(
"train",
format!(
"a leaf of term {features:?} was grown on {cover} rows but only {rows} \
of these rows reach it: pass the rows the model was trained on"
),
));
}
weight.push(if rows > 0.0 {
1.0 / (n_trees * (rows + kappa))
} else {
0.0
});
}
}
let mut part = TermPart {
grid,
trees,
cell_of_row,
weight,
col_mean: Vec::new(),
};
let mut col_sum = vec![0.0; n];
part.add_cell_sums_product(&counts, &mut col_sum);
part.col_mean = col_sum.into_iter().map(|v| v / n as f64).collect();
Ok(part)
}
fn cell_row(&self, cell: usize, scratch: &mut CellScratch) {
let grid = &self.grid;
grid.reset_diff(&mut scratch.diff);
for (t, tree) in self.trees.iter().enumerate() {
let leaf = grid.leaf_of_cell(tree, t, cell);
let w = self.weight[leaf];
if w != 0.0 {
for &b in &grid.leaves[leaf].boxes {
grid.box_add(&mut scratch.diff, b, w);
}
}
}
grid.integrate_into(&scratch.diff, &mut scratch.row);
}
fn add_cell_vector(&self, cell: usize, scratch: &mut CellScratch, out: &mut [f64]) {
self.cell_row(cell, scratch);
for (o, &c) in out.iter_mut().zip(&self.cell_of_row) {
*o += scratch.row[c as usize];
}
}
fn add_product(&self, v: &[f64], out: &mut [f64]) {
let mut sums = vec![0.0; self.grid.len()];
for (&c, &vi) in self.cell_of_row.iter().zip(v) {
sums[c as usize] += vi;
}
self.add_cell_sums_product(&sums, out);
}
fn add_cell_sums_product(&self, sums: &[f64], out: &mut [f64]) {
let grid = &self.grid;
let prefix = grid.prefix(sums);
let mut diff = grid.diff();
for (leaf, &w) in grid.leaves.iter().zip(&self.weight) {
if w == 0.0 {
continue;
}
let total: f64 = leaf.boxes.iter().map(|&b| grid.box_sum(&prefix, b)).sum();
if total != 0.0 {
for &b in &leaf.boxes {
grid.box_add(&mut diff, b, w * total);
}
}
}
let cells = grid.integrate(&diff);
for (o, &c) in out.iter_mut().zip(&self.cell_of_row) {
*o += cells[c as usize];
}
}
pub(super) fn cell_of(&self, data: &DMatrix, row: usize) -> usize {
self.grid.cell_of_row(data, row)
}
}
#[derive(Default)]
struct CellScratch {
diff: Vec<f64>,
row: Vec<f64>,
}
#[derive(Default)]
pub(super) struct TermScratch {
cells: CellScratch,
k: Vec<f64>,
}
pub(super) struct TermKernel<'a> {
n: usize,
pub(super) parts: Vec<TermPart<'a>>,
col_mean: Vec<f64>,
grand_mean: f64,
}
fn center(v: &mut [f64]) {
let mean = v.iter().sum::<f64>() / v.len().max(1) as f64;
for x in v {
*x -= mean;
}
}
impl<'a> TermKernel<'a> {
pub(super) fn new(parts: Vec<TermPart<'a>>, n: usize) -> Self {
let mut col_mean = vec![0.0; n];
for part in &parts {
for (a, &b) in col_mean.iter_mut().zip(&part.col_mean) {
*a += b;
}
}
let grand_mean = col_mean.iter().sum::<f64>() / n.max(1) as f64;
TermKernel {
n,
parts,
col_mean,
grand_mean,
}
}
pub(super) fn add_query(
&self,
t: usize,
cell: usize,
scratch: &mut TermScratch,
out: &mut [f64],
) {
let part = &self.parts[t];
let TermScratch { cells, k } = scratch;
k.clear();
k.resize(self.n, 0.0);
part.add_cell_vector(cell, cells, k);
for (v, &a) in k.iter_mut().zip(&part.col_mean) {
*v -= a;
}
center(k);
for (o, &v) in out.iter_mut().zip(k.iter()) {
*o += v;
}
}
pub(super) fn product(&self, v: &[f64]) -> Vec<f64> {
let mut centered = v.to_vec();
center(&mut centered);
let mut out = vec![0.0; self.n];
for part in &self.parts {
part.add_product(¢ered, &mut out);
}
center(&mut out);
out
}
}
impl Kernel for TermKernel<'_> {
type Scratch = TermScratch;
fn n(&self) -> usize {
self.n
}
fn scratch(&self) -> TermScratch {
TermScratch::default()
}
fn add_row(&self, row: usize, scratch: &mut TermScratch, out: &mut [f64]) {
let TermScratch { cells, k } = scratch;
k.clear();
k.resize(self.n, 0.0);
for part in &self.parts {
part.add_cell_vector(part.cell_of_row[row] as usize, cells, k);
}
let shift = self.grand_mean - self.col_mean[row];
for ((o, &v), &a) in out.iter_mut().zip(k.iter()).zip(&self.col_mean) {
*o += v - a + shift;
}
}
}