use crate::data::DMatrix;
use crate::error::HessboostError;
use crate::tree::regtree::{Node, RegTree};
use rayon::prelude::*;
use serde::{Deserialize, Serialize};
#[derive(Debug, Clone, PartialEq, Serialize, Deserialize, Default)]
#[serde(try_from = "UncheckedLinearLeaves")]
pub struct LinearLeaves {
offsets: Vec<u32>,
intercepts: Vec<f64>,
features: Vec<u32>,
coeffs: Vec<f64>,
}
#[derive(Deserialize)]
pub(crate) struct UncheckedLinearLeaves {
offsets: Vec<u32>,
intercepts: Vec<f64>,
features: Vec<u32>,
coeffs: Vec<f64>,
}
impl UncheckedLinearLeaves {
pub(crate) fn into_unchecked(self) -> LinearLeaves {
LinearLeaves::from_parts(self.offsets, self.intercepts, self.features, self.coeffs)
}
}
impl TryFrom<UncheckedLinearLeaves> for LinearLeaves {
type Error = HessboostError;
fn try_from(unchecked: UncheckedLinearLeaves) -> Result<Self, Self::Error> {
let linear = unchecked.into_unchecked();
if linear.is_consistent() {
Ok(linear)
} else {
Err(HessboostError::model_format(
"linear leaf models are inconsistent",
))
}
}
}
impl LinearLeaves {
#[inline]
pub fn intercept(&self, node: usize) -> f64 {
self.intercepts[node]
}
#[inline]
pub fn terms(&self, node: usize) -> (&[u32], &[f64]) {
let range = self.offsets[node] as usize..self.offsets[node + 1] as usize;
(&self.features[range.clone()], &self.coeffs[range])
}
#[inline]
pub(crate) fn predict(
&self,
node: usize,
constant: f32,
get: impl Fn(u32) -> Option<f32>,
) -> f32 {
let (features, coeffs) = self.terms(node);
let mut out = self.intercepts[node];
for (&f, &c) in features.iter().zip(coeffs) {
match get(f) {
Some(x) => out += c * f64::from(x),
None => return constant,
}
}
out as f32
}
pub(crate) fn scale(&mut self, factor: f64) {
for v in self.intercepts.iter_mut().chain(&mut self.coeffs) {
*v *= factor;
}
}
pub(crate) fn parts(&self) -> (&[u32], &[f64], &[u32], &[f64]) {
(
&self.offsets,
&self.intercepts,
&self.features,
&self.coeffs,
)
}
pub(crate) fn from_parts(
offsets: Vec<u32>,
intercepts: Vec<f64>,
features: Vec<u32>,
coeffs: Vec<f64>,
) -> Self {
LinearLeaves {
offsets,
intercepts,
features,
coeffs,
}
}
fn is_consistent(&self) -> bool {
self.offsets.len() == self.intercepts.len() + 1
&& self.offsets.first() == Some(&0)
&& self
.offsets
.last()
.is_some_and(|&end| end as usize == self.features.len())
&& self.features.len() == self.coeffs.len()
&& self.offsets.windows(2).all(|w| w[0] <= w[1])
&& self
.intercepts
.iter()
.chain(&self.coeffs)
.all(|v| v.is_finite())
}
pub(crate) fn is_valid(&self, nodes: &[Node], n_features: usize) -> bool {
self.is_consistent()
&& self.intercepts.len() == nodes.len()
&& self
.offsets
.windows(2)
.zip(nodes)
.all(|(w, node)| node.is_leaf() || w[0] == w[1])
&& self.features.iter().all(|&f| (f as usize) < n_features)
}
}
pub(crate) fn accumulate_forest(
trees: &[RegTree],
range: std::ops::Range<usize>,
output: impl Fn(usize) -> usize + Sync,
data: &DMatrix,
out: &mut [f32],
k: usize,
weight: impl Fn(usize) -> f32 + Sync,
) {
let row = |(r, out_row): (usize, &mut [f32])| {
for t in range.clone() {
out_row[output(t)] += weight(t) * trees[t].predict_row(data, r);
}
};
if data.n_rows() >= 1024 && rayon::current_num_threads() > 1 {
out.par_chunks_mut(k)
.with_min_len(256)
.enumerate()
.for_each(row);
} else {
out.chunks_mut(k).enumerate().for_each(row);
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn deserialization_refuses_inconsistent_arrays() {
let doc = |offsets: &str| {
format!(r#"{{"offsets":{offsets},"intercepts":[0.5],"features":[0],"coeffs":[2.0]}}"#)
};
let linear: LinearLeaves = serde_json::from_str(&doc("[0,1]")).unwrap();
assert_eq!(linear.terms(0), (&[0u32][..], &[2.0][..]));
for offsets in ["[0,2]", "[1,1]", "[0]", "[0,1,1]"] {
assert!(
serde_json::from_str::<LinearLeaves>(&doc(offsets)).is_err(),
"{offsets}"
);
}
}
}