use crate::data::{DMatrix, FeatureType};
use crate::objective::GradPair;
use crate::tree::builder::LeafRows;
use crate::tree::linear::LinearLeaves;
use crate::tree::regtree::{Node, RegTree};
use rayon::prelude::*;
const ZERO_THRESHOLD: f64 = 1e-35_f32 as f64;
pub(crate) fn fit_linear_leaves(
tree: &mut RegTree,
data: &DMatrix,
gpair: &[GradPair],
rows: &[u32],
lambda: f64,
) {
let nodes = tree.nodes();
if nodes.len() == 1 {
return;
}
let mut members: Vec<Vec<u32>> = vec![Vec::new(); nodes.len()];
let leaves: Vec<u32> = rows
.par_iter()
.with_min_len(4096)
.map(|&r| tree.leaf_id_with(|f| data.get(r as usize, f as usize)) as u32)
.collect();
for (&r, &leaf) in rows.iter().zip(&leaves) {
members[leaf as usize].push(r);
}
fit_leaf_members(tree, data, gpair, &members, lambda);
}
pub(crate) fn fit_captured_linear_leaves(
tree: &mut RegTree,
data: &DMatrix,
gpair: &[GradPair],
leaf_rows: &[LeafRows],
lambda: f64,
) {
let n_nodes = tree.num_nodes();
if n_nodes == 1 {
return;
}
let mut members: Vec<&[u32]> = vec![&[]; n_nodes];
for leaf in leaf_rows {
members[leaf.node] = &leaf.rows;
}
fit_leaf_members(tree, data, gpair, &members, lambda);
}
fn fit_leaf_members(
tree: &mut RegTree,
data: &DMatrix,
gpair: &[GradPair],
members: &[impl AsRef<[u32]> + Sync],
lambda: f64,
) {
let nodes = tree.nodes();
let features = path_features(nodes, data.feature_types());
let models: Vec<(f64, Vec<(u32, f64)>)> = (0..nodes.len())
.into_par_iter()
.map(|id| {
if !nodes[id].is_leaf() {
return (0.0, Vec::new());
}
fit_leaf(&features[id], members[id].as_ref(), data, gpair, lambda)
.unwrap_or((f64::from(nodes[id].leaf_value), Vec::new()))
})
.collect();
let mut offsets = Vec::with_capacity(nodes.len() + 1);
let mut intercepts = Vec::with_capacity(nodes.len());
let mut term_features = Vec::new();
let mut term_coeffs = Vec::new();
offsets.push(0);
for (intercept, terms) in models {
intercepts.push(intercept);
for (f, c) in terms {
term_features.push(f);
term_coeffs.push(c);
}
offsets.push(term_features.len() as u32);
}
tree.set_linear_leaves(LinearLeaves::from_parts(
offsets,
intercepts,
term_features,
term_coeffs,
));
}
fn path_features(nodes: &[Node], types: &[FeatureType]) -> Vec<Vec<u32>> {
let mut out = vec![Vec::new(); nodes.len()];
let mut stack: Vec<(usize, Vec<u32>)> = vec![(0, Vec::new())];
while let Some((id, mut path)) = stack.pop() {
let node = &nodes[id];
if node.is_leaf() {
out[id] = path;
continue;
}
let f = node.split_feature;
let numerical = !node.is_categorical
&& types
.get(f as usize)
.is_none_or(|t| *t == FeatureType::Numerical);
if numerical && let Err(pos) = path.binary_search(&f) {
path.insert(pos, f);
}
stack.push((node.left as usize, path.clone()));
stack.push((node.right as usize, path));
}
out
}
fn fit_leaf(
features: &[u32],
rows: &[u32],
data: &DMatrix,
gpair: &[GradPair],
lambda: f64,
) -> Option<(f64, Vec<(u32, f64)>)> {
let p = features.len();
let dim = p + 1;
let mut xthx = vec![0.0f64; dim * (dim + 1) / 2];
let mut xtg = vec![0.0f64; dim];
let mut x = vec![0.0f32; dim];
x[p] = 1.0;
let mut complete = 0usize;
'rows: for &r in rows {
for (slot, &f) in x.iter_mut().zip(features) {
match data.get(r as usize, f as usize) {
Some(v) => *slot = v,
None => continue 'rows,
}
}
complete += 1;
let GradPair { grad, hess } = gpair[r as usize];
let mut k = 0;
for i in 0..dim {
let xi = f64::from(x[i]);
xtg[i] += xi * f64::from(grad);
let xih = xi * f64::from(hess);
for &xj in &x[i..] {
xthx[k] += xih * f64::from(xj);
k += 1;
}
}
}
if complete < dim {
return None;
}
let mut a = vec![0.0f64; dim * dim];
let mut k = 0;
for i in 0..dim {
for j in i..dim {
a[i * dim + j] = xthx[k];
a[j * dim + i] = xthx[k];
k += 1;
}
if i < p {
a[i * dim + i] += lambda;
}
}
let rhs: Vec<f64> = xtg.iter().map(|v| -v).collect();
let beta = solve_full_pivot(a, rhs)?;
let terms = features
.iter()
.zip(&beta)
.filter(|&(_, c)| c.abs() > ZERO_THRESHOLD)
.map(|(&f, &c)| (f, c))
.collect();
Some((beta[p], terms))
}
fn solve_full_pivot(mut a: Vec<f64>, mut b: Vec<f64>) -> Option<Vec<f64>> {
let n = b.len();
let mut cols: Vec<usize> = (0..n).collect();
let mut first_pivot = 0.0f64;
for k in 0..n {
let (mut pr, mut pc, mut pv) = (k, k, 0.0f64);
for r in k..n {
for c in k..n {
let v = a[r * n + c].abs();
if v > pv {
(pr, pc, pv) = (r, c, v);
}
}
}
if k == 0 {
first_pivot = pv;
}
if pv <= first_pivot * f64::EPSILON * n as f64 {
return None;
}
if pr != k {
for c in 0..n {
a.swap(pr * n + c, k * n + c);
}
b.swap(pr, k);
}
if pc != k {
for r in 0..n {
a.swap(r * n + pc, r * n + k);
}
cols.swap(pc, k);
}
let pivot = a[k * n + k];
for r in k + 1..n {
let factor = a[r * n + k] / pivot;
if factor == 0.0 {
continue;
}
for c in k..n {
a[r * n + c] -= factor * a[k * n + c];
}
b[r] -= factor * b[k];
}
}
let mut y = vec![0.0f64; n];
for k in (0..n).rev() {
let mut s = b[k];
for c in k + 1..n {
s -= a[k * n + c] * y[c];
}
y[k] = s / a[k * n + k];
}
let mut x = vec![0.0f64; n];
for (j, &unknown) in cols.iter().enumerate() {
x[unknown] = y[j];
}
x.iter().all(|v| v.is_finite()).then_some(x)
}
#[cfg(test)]
mod tests {
use super::*;
use crate::tree::{ChildLeaf, SplitRule};
#[test]
fn full_pivot_solver_matches_a_known_solution() {
let a = vec![1.0, 2.0, 3.0, 2.0, 5.0, 3.0, 1.0, 0.0, 8.0];
let want = [-40.0, 16.0, 5.0];
let b: Vec<f64> = (0..3)
.map(|r| (0..3).map(|c| a[r * 3 + c] * want[c]).sum())
.collect();
let got = solve_full_pivot(a, b).unwrap();
for (g, w) in got.iter().zip(want) {
assert!((g - w).abs() < 1e-9, "{got:?}");
}
assert!(solve_full_pivot(vec![1.0, 2.0, 2.0, 4.0], vec![1.0, 2.0]).is_none());
}
fn stump(default_left: bool) -> RegTree {
let mut tree = RegTree::with_root(1.0);
tree.expand(
0,
SplitRule::numeric(0, 0.0, default_left),
ChildLeaf::new(-0.5, 1.0),
ChildLeaf::new(0.25, 1.0),
);
tree
}
fn column(x: &[f32]) -> DMatrix {
DMatrix::from_dense(x, x.len(), 1).unwrap()
}
fn line_gradients(x: &[f32]) -> Vec<GradPair> {
x.iter()
.map(|&v| GradPair::new(-(2.0 * v + 1.0), 1.0))
.collect()
}
#[test]
fn leaf_fit_is_the_ridge_newton_step() {
let x = [-1.0f32, 1.0, 2.0, 3.0, 4.0];
let data = column(&x);
let gpair = line_gradients(&x);
let rows = [0, 1, 2, 3, 4];
let lambda = 2.0;
let mut tree = stump(true);
fit_linear_leaves(&mut tree, &data, &gpair, &rows, lambda);
let (sxx, sx, n, sxy, sy) = (30.0 + lambda, 10.0, 4.0, 70.0, 24.0);
let det = sxx * n - sx * sx;
let slope = (n * sxy - sx * sy) / det;
let intercept = (sxx * sy - sx * sxy) / det;
let linear = tree.linear_leaves().unwrap();
let (features, coeffs) = linear.terms(2);
assert_eq!(features, &[0]);
assert!((coeffs[0] - slope).abs() < 1e-12, "{coeffs:?} vs {slope}");
assert!((linear.intercept(2) - intercept).abs() < 1e-12);
assert!(linear.terms(1).0.is_empty());
assert_eq!(tree.predict_row(&data, 0), -0.5);
let mut exact = stump(true);
fit_linear_leaves(&mut exact, &data, &gpair, &rows, 0.0);
for (r, &v) in x.iter().enumerate().skip(1) {
assert!((exact.predict_row(&data, r) - (2.0 * v + 1.0)).abs() < 1e-5);
}
}
#[test]
fn missing_features_are_skipped_in_the_fit_and_predict_the_constant() {
let x = [-1.0f32, 1.0, 2.0, f32::NAN, 3.0, 4.0];
let data = column(&x);
let mut gpair = line_gradients(&x);
gpair[3] = GradPair::new(1e6, 1.0);
let mut tree = stump(false);
fit_linear_leaves(&mut tree, &data, &gpair, &[0, 1, 2, 3, 4, 5], 0.0);
assert_eq!(tree.predict_row(&data, 3), 0.25);
for r in [1, 2, 4, 5] {
let want = 2.0 * x[r] + 1.0;
assert!((tree.predict_row(&data, r) - want).abs() < 1e-5, "row {r}");
}
}
#[test]
fn categorical_path_features_route_but_stay_out_of_the_model() {
let mut tree = RegTree::with_root(8.0);
let (l, r) = tree.expand(
0,
SplitRule::categorical(0, &[1], false),
ChildLeaf::new(0.0, 4.0),
ChildLeaf::new(0.0, 4.0),
);
for child in [l, r] {
tree.expand(
child,
SplitRule::numeric(1, 0.5, true),
ChildLeaf::new(0.0, 2.0),
ChildLeaf::new(0.0, 2.0),
);
}
let types = [FeatureType::Categorical, FeatureType::Numerical];
let paths = path_features(tree.nodes(), &types);
for (id, node) in tree.nodes().iter().enumerate() {
let want: &[u32] = if node.is_leaf() { &[1] } else { &[] };
assert_eq!(paths[id], want, "node {id}");
}
}
}