qualia_core_db/solvers/learning/trees/
boosting.rs1use crate::solvers::learning::trees::decision_tree::{DecisionTree, TreeParams};
7use crate::solvers::learning::LearningError;
8use crate::solvers::statistics::descriptive::mean;
9
10#[derive(Debug, Clone)]
12pub struct GradientBoosting {
13 init: f64,
14 trees: Vec<DecisionTree>,
15 learning_rate: f64,
16 p: usize,
17}
18
19impl GradientBoosting {
20 pub fn fit_regressor(
24 x: &[f64],
25 y: &[f64],
26 n: usize,
27 p: usize,
28 n_estimators: usize,
29 learning_rate: f64,
30 params: TreeParams,
31 ) -> Result<Self, LearningError> {
32 if n == 0 || p == 0 || x.len() != n * p || y.len() != n {
33 return Err(LearningError::InvalidDimension);
34 }
35 if n_estimators == 0 || !(learning_rate > 0.0) {
36 return Err(LearningError::InsufficientData);
37 }
38
39 let init = mean(y).ok_or(LearningError::InsufficientData)?;
40 let mut pred = vec![init; n];
41 let mut residual = vec![0.0; n];
42 let mut trees = Vec::with_capacity(n_estimators);
43
44 for _ in 0..n_estimators {
45 for i in 0..n {
46 residual[i] = y[i] - pred[i]; }
48 let tree = DecisionTree::fit_regressor(x, &residual, n, p, params)?;
49 for i in 0..n {
51 pred[i] += learning_rate * tree.predict_row(&x[i * p..(i + 1) * p]);
52 }
53 trees.push(tree);
54 }
55
56 Ok(Self {
57 init,
58 trees,
59 learning_rate,
60 p,
61 })
62 }
63
64 pub fn predict_row(&self, q: &[f64]) -> f64 {
65 self.init + self.learning_rate * self.trees.iter().map(|t| t.predict_row(q)).sum::<f64>()
66 }
67
68 pub fn predict(&self, x: &[f64], m: usize) -> Vec<f64> {
69 (0..m)
70 .map(|i| self.predict_row(&x[i * self.p..(i + 1) * self.p]))
71 .collect()
72 }
73
74 pub fn n_estimators(&self) -> usize {
75 self.trees.len()
76 }
77}
78
79#[cfg(test)]
80mod tests {
81 use super::*;
82 use crate::solvers::learning::metrics::regression::{mse, r2_score};
83
84 #[test]
85 fn boosting_reduces_error_with_more_stages() {
86 let n = 50;
88 let x: Vec<f64> = (0..n).map(|i| i as f64 / 10.0).collect();
89 let y: Vec<f64> = x.iter().map(|&xi| (xi).sin() * 3.0 + 0.5 * xi).collect();
90 let params = TreeParams {
91 max_depth: 3,
92 ..TreeParams::default()
93 };
94 let few = GradientBoosting::fit_regressor(&x, &y, n, 1, 5, 0.1, params).unwrap();
95 let many = GradientBoosting::fit_regressor(&x, &y, n, 1, 200, 0.1, params).unwrap();
96 let mse_few = mse(&y, &few.predict(&x, n)).unwrap();
97 let mse_many = mse(&y, &many.predict(&x, n)).unwrap();
98 assert!(
99 mse_many < mse_few,
100 "more stages should fit better: {mse_many} !< {mse_few}"
101 );
102 assert!(r2_score(&y, &many.predict(&x, n)).unwrap() > 0.9);
104 }
105
106 #[test]
107 fn single_stage_is_init_plus_one_tree() {
108 let x = [1.0, 2.0, 3.0, 4.0, 5.0, 6.0];
109 let y = [1.0, 2.0, 1.0, 5.0, 6.0, 5.0];
110 let gb = GradientBoosting::fit_regressor(
111 &x,
112 &y,
113 6,
114 1,
115 1,
116 0.5,
117 TreeParams {
118 max_depth: 2,
119 ..TreeParams::default()
120 },
121 )
122 .unwrap();
123 assert_eq!(gb.n_estimators(), 1);
124 let p = gb.predict_row(&[3.5]);
126 assert!(p.is_finite());
127 }
128
129 #[test]
130 fn guards() {
131 assert_eq!(
132 GradientBoosting::fit_regressor(&[1.0], &[1.0], 1, 1, 0, 0.1, TreeParams::default())
133 .unwrap_err(),
134 LearningError::InsufficientData
135 );
136 assert_eq!(
137 GradientBoosting::fit_regressor(&[1.0], &[1.0], 1, 1, 10, 0.0, TreeParams::default())
138 .unwrap_err(),
139 LearningError::InsufficientData
140 );
141 }
142}