use rayon::prelude::*;
use crate::error::{HessboostError, Result};
use crate::model::Predictions;
use crate::tree::RegTree;
pub(super) trait Kernel: Sync {
type Scratch: Send;
fn n(&self) -> usize;
fn scratch(&self) -> Self::Scratch;
fn add_row(&self, row: usize, scratch: &mut Self::Scratch, out: &mut [f64]);
fn dense(&self) -> Vec<f64> {
let n = self.n();
let mut k = vec![0.0; n * n];
k.par_chunks_mut(n.max(1)).enumerate().for_each_init(
|| self.scratch(),
|scratch, (row, out)| self.add_row(row, scratch, out),
);
k
}
}
const INTERNAL: u32 = u32::MAX;
pub(super) struct LeafKernel {
n: usize,
n_trees: usize,
node_offset: Vec<usize>,
node_leaf: Vec<u32>,
leaf_weight: Vec<f64>,
leaf_start: Vec<usize>,
leaf_rows: Vec<u32>,
row_leaves: Vec<u32>,
}
impl LeafKernel {
pub(super) fn new(trees: &[RegTree], node_ids: &Predictions<u32>, kappa: f64) -> Result<Self> {
let n_trees = trees.len();
debug_assert_eq!(node_ids.width(), n_trees);
let n = node_ids.n_rows();
let mut node_offset = Vec::with_capacity(n_trees + 1);
let mut node_leaf = Vec::new();
let mut n_leaves = 0u32;
for tree in trees {
node_offset.push(node_leaf.len());
for node in tree.nodes() {
if node.is_leaf() {
node_leaf.push(n_leaves);
n_leaves += 1;
} else {
node_leaf.push(INTERNAL);
}
}
}
node_offset.push(node_leaf.len());
let mut row_leaves = vec![0u32; n * n_trees];
let mut counts = vec![0usize; n_leaves as usize];
for (row, ids) in node_ids.rows().enumerate() {
for (t, &id) in ids.iter().enumerate() {
let leaf = node_leaf[node_offset[t] + id as usize];
row_leaves[row * n_trees + t] = leaf;
counts[leaf as usize] += 1;
}
}
for (t, tree) in trees.iter().enumerate() {
for (id, node) in tree.nodes().iter().enumerate() {
if !node.is_leaf() {
continue;
}
let count = counts[node_leaf[node_offset[t] + id] as usize];
if (count as f64) < f64::from(node.sum_hess) {
return Err(HessboostError::invalid_data(
"train",
format!(
"leaf {id} of tree {t} was grown on {} rows but only {count} of these \
rows reach it: pass the rows the model was trained on (or refit on)",
node.sum_hess
),
));
}
}
}
let scale = n_trees as f64;
let leaf_weight = counts
.iter()
.map(|&c| {
if c == 0 {
0.0
} else {
1.0 / (scale * (c as f64 + kappa))
}
})
.collect();
let mut leaf_start = Vec::with_capacity(counts.len() + 1);
let mut at = 0usize;
for &c in &counts {
leaf_start.push(at);
at += c;
}
leaf_start.push(at);
let mut fill = leaf_start.clone();
let mut leaf_rows = vec![0u32; at];
for (row, leaves) in row_leaves.chunks_exact(n_trees.max(1)).enumerate() {
for &leaf in leaves {
let slot = &mut fill[leaf as usize];
leaf_rows[*slot] = row as u32;
*slot += 1;
}
}
Ok(LeafKernel {
n,
n_trees,
node_offset,
node_leaf,
leaf_weight,
leaf_start,
leaf_rows,
row_leaves,
})
}
pub(super) fn n_trees(&self) -> usize {
self.n_trees
}
fn leaf_of(&self, t: usize, id: u32) -> u32 {
self.node_leaf[self.node_offset[t] + id as usize]
}
pub(super) fn add_query(&self, node_ids: &[u32], out: &mut [f64]) {
for (t, &id) in node_ids.iter().enumerate() {
self.add_leaf(self.leaf_of(t, id), out);
}
}
fn add_leaf(&self, leaf: u32, out: &mut [f64]) {
let leaf = leaf as usize;
let w = self.leaf_weight[leaf];
for &row in &self.leaf_rows[self.leaf_start[leaf]..self.leaf_start[leaf + 1]] {
out[row as usize] += w;
}
}
}
impl Kernel for LeafKernel {
type Scratch = ();
fn n(&self) -> usize {
self.n
}
fn scratch(&self) {}
fn add_row(&self, row: usize, (): &mut (), out: &mut [f64]) {
let leaves = &self.row_leaves[row * self.n_trees..(row + 1) * self.n_trees];
for &leaf in leaves {
self.add_leaf(leaf, out);
}
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::tree::{ChildLeaf, SplitRule};
#[test]
fn kernel_rows_are_the_leaf_matrix_rows() {
let empty = LeafKernel::new(&[], &Predictions::new(Vec::new(), 5, 0), 0.0).unwrap();
assert_eq!((empty.n(), empty.n_trees()), (5, 0));
let mut tree = RegTree::with_root(0.0);
let (left, right) = tree.expand(
0,
SplitRule::numeric(0, 0.5, true),
ChildLeaf::new(0.0, 2.0),
ChildLeaf::new(0.0, 1.0),
);
let ids = Predictions::new(vec![left as u32, left as u32, right as u32], 3, 1);
let kernel = LeafKernel::new(std::slice::from_ref(&tree), &ids, 0.0).unwrap();
assert_eq!(kernel.n(), 3);
assert_eq!(
kernel.dense(),
vec![0.5, 0.5, 0.0, 0.5, 0.5, 0.0, 0.0, 0.0, 1.0]
);
}
}