Skip to main content

qualia_core_db/solvers/learning/trees/
decision_tree.rs

1//! CART decision trees (ISL ch 8.1) — recursive binary splitting for regression
2//! (variance / MSE reduction) and classification (Gini impurity), over a row-major
3//! feature matrix. The arena-based node store avoids deep `Box` recursion. The same
4//! builder powers random forests and gradient boosting (with feature subsampling /
5//! shallow depth). Scalar split search → CPU.
6
7use crate::solvers::learning::LearningError;
8
9/// Split criterion.
10#[derive(Debug, Clone, Copy, PartialEq, Eq)]
11pub enum Criterion {
12    /// Regression: minimise within-node sum of squared error.
13    Mse,
14    /// Classification: minimise Gini impurity (labels carried as integers-in-f64).
15    Gini,
16}
17
18/// Tree hyper-parameters.
19#[derive(Debug, Clone, Copy)]
20pub struct TreeParams {
21    pub max_depth: usize,
22    pub min_samples_split: usize,
23    pub min_samples_leaf: usize,
24    /// Features considered per split (`None` = all). `Some(m)` random-samples `m`
25    /// features — the mechanism random forests use.
26    pub max_features: Option<usize>,
27    pub seed: u64,
28}
29
30impl Default for TreeParams {
31    fn default() -> Self {
32        Self {
33            max_depth: 8,
34            min_samples_split: 2,
35            min_samples_leaf: 1,
36            max_features: None,
37            seed: 0,
38        }
39    }
40}
41
42#[derive(Debug, Clone)]
43struct Node {
44    feature: i32, // -1 ⇒ leaf
45    threshold: f64,
46    left: usize,
47    right: usize,
48    value: f64, // leaf prediction (regression mean / majority class label)
49}
50
51/// A fitted decision tree.
52#[derive(Debug, Clone)]
53pub struct DecisionTree {
54    nodes: Vec<Node>,
55    criterion: Criterion,
56    p: usize,
57}
58
59struct Lcg(u64);
60impl Lcg {
61    fn next(&mut self) -> u64 {
62        self.0 = self
63            .0
64            .wrapping_mul(6364136223846793005)
65            .wrapping_add(1442695040888963407);
66        self.0
67    }
68}
69
70/// Impurity of the y-values at the indices, per criterion.
71fn impurity(y: &[f64], idx: &[usize], criterion: Criterion) -> f64 {
72    let n = idx.len();
73    if n == 0 {
74        return 0.0;
75    }
76    match criterion {
77        Criterion::Mse => {
78            let mean = idx.iter().map(|&i| y[i]).sum::<f64>() / n as f64;
79            idx.iter()
80                .map(|&i| (y[i] - mean) * (y[i] - mean))
81                .sum::<f64>()
82                / n as f64
83        }
84        Criterion::Gini => {
85            // Class labels are small non-negative integers stored as f64.
86            let max_label = idx.iter().map(|&i| y[i] as usize).max().unwrap_or(0);
87            let mut counts = vec![0usize; max_label + 1];
88            for &i in idx {
89                counts[y[i] as usize] += 1;
90            }
91            let nf = n as f64;
92            1.0 - counts
93                .iter()
94                .map(|&c| {
95                    let p = c as f64 / nf;
96                    p * p
97                })
98                .sum::<f64>()
99        }
100    }
101}
102
103/// Leaf prediction for the y-values at the indices.
104fn leaf_value(y: &[f64], idx: &[usize], criterion: Criterion) -> f64 {
105    match criterion {
106        Criterion::Mse => idx.iter().map(|&i| y[i]).sum::<f64>() / idx.len() as f64,
107        Criterion::Gini => {
108            let max_label = idx.iter().map(|&i| y[i] as usize).max().unwrap_or(0);
109            let mut counts = vec![0usize; max_label + 1];
110            for &i in idx {
111                counts[y[i] as usize] += 1;
112            }
113            counts
114                .iter()
115                .enumerate()
116                .max_by_key(|(_, &c)| c)
117                .map(|(l, _)| l)
118                .unwrap_or(0) as f64
119        }
120    }
121}
122
123struct Builder<'a> {
124    x: &'a [f64],
125    y: &'a [f64],
126    p: usize,
127    params: TreeParams,
128    criterion: Criterion,
129    nodes: Vec<Node>,
130    rng: Lcg,
131}
132
133impl<'a> Builder<'a> {
134    fn build(&mut self, idx: &[usize], depth: usize) -> usize {
135        let value = leaf_value(self.y, idx, self.criterion);
136        let node_impurity = impurity(self.y, idx, self.criterion);
137        // Stopping rules.
138        if depth >= self.params.max_depth
139            || idx.len() < self.params.min_samples_split
140            || node_impurity <= 1e-12
141        {
142            return self.push_leaf(value);
143        }
144
145        // Candidate feature set (optionally subsampled — random forests).
146        let feats = self.candidate_features();
147
148        let mut best_gain = 0.0;
149        let mut best_feat = usize::MAX;
150        let mut best_thr = 0.0;
151        let parent = node_impurity;
152        let n = idx.len() as f64;
153
154        for &f in &feats {
155            // Sort the node's rows by this feature.
156            let mut order: Vec<usize> = idx.to_vec();
157            order.sort_by(|&a, &b| {
158                self.x[a * self.p + f]
159                    .partial_cmp(&self.x[b * self.p + f])
160                    .unwrap_or(core::cmp::Ordering::Equal)
161            });
162            // Try thresholds between distinct consecutive values.
163            for s in 1..order.len() {
164                let v0 = self.x[order[s - 1] * self.p + f];
165                let v1 = self.x[order[s] * self.p + f];
166                if v1 <= v0 {
167                    continue; // identical values — no split here
168                }
169                let (left, right) = order.split_at(s);
170                if left.len() < self.params.min_samples_leaf
171                    || right.len() < self.params.min_samples_leaf
172                {
173                    continue;
174                }
175                let il = impurity(self.y, left, self.criterion);
176                let ir = impurity(self.y, right, self.criterion);
177                let gain = parent - (left.len() as f64 / n * il + right.len() as f64 / n * ir);
178                if gain > best_gain {
179                    best_gain = gain;
180                    best_feat = f;
181                    best_thr = 0.5 * (v0 + v1);
182                }
183            }
184        }
185
186        if best_feat == usize::MAX || best_gain <= 1e-12 {
187            return self.push_leaf(value);
188        }
189
190        // Partition and recurse.
191        let mut left_idx = Vec::new();
192        let mut right_idx = Vec::new();
193        for &i in idx {
194            if self.x[i * self.p + best_feat] <= best_thr {
195                left_idx.push(i);
196            } else {
197                right_idx.push(i);
198            }
199        }
200        // Reserve this node's slot before recursing (children get later indices).
201        let me = self.nodes.len();
202        self.nodes.push(Node {
203            feature: best_feat as i32,
204            threshold: best_thr,
205            left: 0,
206            right: 0,
207            value,
208        });
209        let l = self.build(&left_idx, depth + 1);
210        let r = self.build(&right_idx, depth + 1);
211        self.nodes[me].left = l;
212        self.nodes[me].right = r;
213        me
214    }
215
216    fn candidate_features(&mut self) -> Vec<usize> {
217        match self.params.max_features {
218            None => (0..self.p).collect(),
219            Some(m) => {
220                let m = m.clamp(1, self.p);
221                let mut feats: Vec<usize> = (0..self.p).collect();
222                // Partial Fisher–Yates to pick m features.
223                for i in 0..m {
224                    let j = i + (self.rng.next() as usize) % (self.p - i);
225                    feats.swap(i, j);
226                }
227                feats.truncate(m);
228                feats
229            }
230        }
231    }
232
233    fn push_leaf(&mut self, value: f64) -> usize {
234        let idx = self.nodes.len();
235        self.nodes.push(Node {
236            feature: -1,
237            threshold: 0.0,
238            left: 0,
239            right: 0,
240            value,
241        });
242        idx
243    }
244}
245
246impl DecisionTree {
247    fn fit_inner(
248        x: &[f64],
249        y: &[f64],
250        n: usize,
251        p: usize,
252        criterion: Criterion,
253        params: TreeParams,
254    ) -> Result<Self, LearningError> {
255        if n == 0 || p == 0 || x.len() != n * p || y.len() != n {
256            return Err(LearningError::InvalidDimension);
257        }
258        let mut builder = Builder {
259            x,
260            y,
261            p,
262            params,
263            criterion,
264            nodes: Vec::new(),
265            rng: Lcg(params.seed ^ 0x9E3779B97F4A7C15),
266        };
267        let idx: Vec<usize> = (0..n).collect();
268        builder.build(&idx, 0);
269        Ok(Self {
270            nodes: builder.nodes,
271            criterion,
272            p,
273        })
274    }
275
276    /// Fit a regression tree (`Criterion::Mse`).
277    pub fn fit_regressor(
278        x: &[f64],
279        y: &[f64],
280        n: usize,
281        p: usize,
282        params: TreeParams,
283    ) -> Result<Self, LearningError> {
284        Self::fit_inner(x, y, n, p, Criterion::Mse, params)
285    }
286
287    /// Fit a classification tree (`Criterion::Gini`); labels are small integers.
288    pub fn fit_classifier(
289        x: &[f64],
290        y: &[usize],
291        n: usize,
292        p: usize,
293        params: TreeParams,
294    ) -> Result<Self, LearningError> {
295        let yf: Vec<f64> = y.iter().map(|&v| v as f64).collect();
296        Self::fit_inner(x, &yf, n, p, Criterion::Gini, params)
297    }
298
299    /// Raw leaf prediction (regression value, or class label as f64) for one row.
300    pub fn predict_row(&self, q: &[f64]) -> f64 {
301        let mut node = 0;
302        loop {
303            let nd = &self.nodes[node];
304            if nd.feature < 0 {
305                return nd.value;
306            }
307            node = if q[nd.feature as usize] <= nd.threshold {
308                nd.left
309            } else {
310                nd.right
311            };
312        }
313    }
314
315    pub fn predict(&self, x: &[f64], m: usize) -> Vec<f64> {
316        (0..m)
317            .map(|i| self.predict_row(&x[i * self.p..(i + 1) * self.p]))
318            .collect()
319    }
320
321    /// Class label prediction (rounds the leaf value) for a classifier tree.
322    pub fn predict_class(&self, q: &[f64]) -> usize {
323        self.predict_row(q).round() as usize
324    }
325
326    pub fn criterion(&self) -> Criterion {
327        self.criterion
328    }
329}
330
331#[cfg(test)]
332mod tests {
333    use super::*;
334
335    #[test]
336    fn regression_tree_fits_a_step_function() {
337        // y is a step: 0 for x<5, 10 for x>=5. A depth-1 tree captures it exactly.
338        let x: Vec<f64> = (0..10).map(|i| i as f64).collect();
339        let y: Vec<f64> = x
340            .iter()
341            .map(|&xi| if xi < 5.0 { 0.0 } else { 10.0 })
342            .collect();
343        let t = DecisionTree::fit_regressor(&x, &y, 10, 1, TreeParams::default()).unwrap();
344        assert!((t.predict_row(&[2.0]) - 0.0).abs() < 1e-9);
345        assert!((t.predict_row(&[8.0]) - 10.0).abs() < 1e-9);
346    }
347
348    #[test]
349    fn classification_tree_separates_xor_free_data() {
350        // Two features; class determined by x0 threshold.
351        let x = [0.0, 0.0, 1.0, 1.0, 8.0, 0.0, 9.0, 1.0];
352        let y = [0usize, 0, 1, 1];
353        let t = DecisionTree::fit_classifier(&x, &y, 4, 2, TreeParams::default()).unwrap();
354        assert_eq!(t.predict_class(&[0.5, 0.5]), 0);
355        assert_eq!(t.predict_class(&[8.5, 0.5]), 1);
356    }
357
358    #[test]
359    fn perfectly_fits_training_set_when_deep() {
360        // Distinct x → a deep tree memorizes the targets (train error 0).
361        let x = [1.0, 2.0, 3.0, 4.0, 5.0];
362        let y = [3.0, 1.0, 4.0, 1.0, 5.0];
363        let t = DecisionTree::fit_regressor(&x, &y, 5, 1, TreeParams::default()).unwrap();
364        for i in 0..5 {
365            assert!((t.predict_row(&[x[i]]) - y[i]).abs() < 1e-9);
366        }
367    }
368
369    #[test]
370    fn depth_limit_makes_a_single_leaf() {
371        let x = [1.0, 2.0, 3.0, 4.0];
372        let y = [1.0, 2.0, 3.0, 4.0];
373        let params = TreeParams {
374            max_depth: 0,
375            ..TreeParams::default()
376        };
377        let t = DecisionTree::fit_regressor(&x, &y, 4, 1, params).unwrap();
378        // A single leaf predicts the global mean (2.5) for everything.
379        assert!((t.predict_row(&[1.0]) - 2.5).abs() < 1e-9);
380        assert!((t.predict_row(&[4.0]) - 2.5).abs() < 1e-9);
381    }
382}