qualia_core_db/solvers/learning/trees/
decision_tree.rs1use crate::solvers::learning::LearningError;
8
9#[derive(Debug, Clone, Copy, PartialEq, Eq)]
11pub enum Criterion {
12 Mse,
14 Gini,
16}
17
18#[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 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, threshold: f64,
46 left: usize,
47 right: usize,
48 value: f64, }
50
51#[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
70fn 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 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
103fn 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 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 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 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 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; }
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 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 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 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 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 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 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 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 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 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 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 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}