use serde::{Deserialize, Serialize};
use crate::data::DMatrix;
use crate::error::{HessboostError, Result};
use super::initial_margins;
pub(crate) fn shrink_margins(margins: &mut [f32], factor: f64) {
if factor != 1.0 {
for m in margins {
*m = (f64::from(*m) * factor) as f32;
}
}
}
#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
pub(crate) struct Shrinkage {
factors: Vec<f64>,
base_score: Vec<f32>,
}
impl Shrinkage {
pub(crate) fn new(factors: Vec<f64>, base_score: Vec<f32>) -> Self {
Shrinkage {
factors,
base_score,
}
}
pub(crate) fn factors(&self) -> &[f64] {
&self.factors
}
pub(crate) fn base_score(&self) -> &[f32] {
&self.base_score
}
pub(crate) fn truncated(&self, k: usize) -> Shrinkage {
Shrinkage::new(self.factors[..k].to_vec(), self.base_score.clone())
}
pub(crate) fn scaling(&self, k: usize, trees_per_iteration: usize) -> (Vec<f32>, Vec<f32>) {
let mut weights = vec![0.0f32; k * trees_per_iteration];
let mut product = 1.0f64;
for (i, layer) in weights
.chunks_exact_mut(trees_per_iteration)
.enumerate()
.rev()
{
layer.fill(product as f32);
product *= self.factors[i];
}
let base = self
.base_score
.iter()
.map(|&b| (f64::from(b) * product) as f32)
.collect();
(weights, base)
}
pub(crate) fn matches(
&self,
trees_per_iteration: usize,
tree_weight: impl Fn(usize) -> f32,
base_score: &[f32],
) -> bool {
let (weights, base) = self.scaling(self.factors.len(), trees_per_iteration);
let same = |a: f32, b: f32| a.to_bits() == b.to_bits();
weights
.iter()
.enumerate()
.all(|(t, &w)| same(w, tree_weight(t)))
&& base.len() == base_score.len()
&& base.iter().zip(base_score).all(|(&a, &b)| same(a, b))
}
pub(crate) fn start_margins(&self, data: &DMatrix) -> Vec<f32> {
if data.base_margin().is_some() {
vec![0.0; data.n_rows() * self.base_score.len()]
} else {
initial_margins(&self.base_score, data)
}
}
pub(crate) fn finish_margins(&self, data: &DMatrix, margins: &mut [f32]) {
if data.base_margin().is_some() {
let base = initial_margins(&self.base_score, data);
for (m, b) in margins.iter_mut().zip(base) {
*m += b;
}
}
}
pub(crate) fn validate(&self, iterations: usize, n_outputs: usize) -> Result<()> {
if self.factors.len() != iterations {
return Err(HessboostError::ModelFormat(format!(
"the shrinkage record holds {} coefficients for {iterations} iterations",
self.factors.len()
)));
}
if self
.factors
.iter()
.any(|&s| !(s.is_finite() && s > 0.0 && s <= 1.0))
|| self.factors.first().is_some_and(|&s| s != 1.0)
{
return Err(HessboostError::model_format(
"shrinkage coefficients must be in (0, 1], the first one 1",
));
}
if self.base_score.len() != n_outputs || self.base_score.iter().any(|b| !b.is_finite()) {
return Err(HessboostError::ModelFormat(format!(
"the shrinkage record must hold one finite intercept per output ({n_outputs} \
outputs, got {:?})",
self.base_score
)));
}
Ok(())
}
}