From 41ca9f26eb8ab45ce62d3a658528baf39eadc5ec Mon Sep 17 00:00:00 2001 From: slievens <5492873+slievens@users.noreply.github.com> Date: Fri, 2 Oct 2026 11:24:12 +0200 Subject: [PATCH 01/17] First step in making rf faster --- src/tree/base_tree_regressor.rs | 266 ++++++++++++++++++-------------- 1 file changed, 146 insertions(+), 120 deletions(-) diff --git a/src/tree/base_tree_regressor.rs b/src/tree/base_tree_regressor.rs index b833e345..0fe54bae 100644 --- a/src/tree/base_tree_regressor.rs +++ b/src/tree/base_tree_regressor.rs @@ -130,9 +130,8 @@ struct NodeVisitor<'a, TX: Number + PartialOrd, TY: Number, X: Array2, Y: Ar x: &'a X, y: &'a Y, node: usize, - samples: Vec, - sample_weights: Option<&'a [f64]>, - order: &'a [Vec], + // holds the elements for this node, sorted for each feature [num_features, num_samples] + sorted_node_elements: Vec>, true_child_output: f64, false_child_output: f64, level: u16, @@ -145,9 +144,7 @@ impl<'a, TX: Number + PartialOrd, TY: Number, X: Array2, Y: Array1> { fn new( node_id: usize, - samples: Vec, - sample_weights: Option<&'a [f64]>, - order: &'a [Vec], + sorted_node_elements: Vec>, x: &'a X, y: &'a Y, level: u16, @@ -156,9 +153,7 @@ impl<'a, TX: Number + PartialOrd, TY: Number, X: Array2, Y: Array1> x, y, node: node_id, - samples, - sample_weights, - order, + sorted_node_elements, true_child_output: 0f64, false_child_output: 0f64, level, @@ -167,9 +162,12 @@ impl<'a, TX: Number + PartialOrd, TY: Number, X: Array2, Y: Array1> } } - /// Weighted count of sample `i`. The weight is 1.0 if no weights are given. - fn mass_of(&self, i: usize) -> f64 { - mass_of(i, &self.samples, self.sample_weights) + /// number of samples in this node (nodevisitor) + fn num_samples(&self) -> usize { + self.sorted_node_elements[0] + .iter() + .map(|node_elem| node_elem.count) + .sum() } } @@ -181,6 +179,30 @@ fn mass_of(i: usize, samples: &[usize], sample_weights: Option<&[f64]>) -> f64 { } } +// Struct representing an element that belongs logically to a Node, as stored in NodeVisitor +#[derive(Copy, Clone)] +struct NodeElement { + // the row index in the dataset + pub row_idx: usize, + // the number of times this row is present, should always be > 0 + pub count: usize, + // total mass of this element, equals count * mass of individual element. + // equals count when no sample weights were used + pub mass: f64, +} + +// Checks whether the example indicated by node_element is a "true child" for this node +fn is_true_sample(node_element: &NodeElement, x: &X, node: &Node) -> bool +where + TX: Number + PartialOrd, + X: Array2, +{ + x.get((node_element.row_idx, node.split_feature)) + .to_f64() + .unwrap() + <= node.split_value.unwrap_or(f64::NAN) +} + impl, Y: Array1> BaseTreeRegressor { @@ -246,6 +268,21 @@ impl, Y: Array1> order.push(col_i.argsort_mut()); } + let sorted_node_elements: Vec> = order + .iter() + .map(|col_order| { + col_order + .iter() + .filter(|&&i| samples[i] > 0) + .map(|&i| NodeElement { + row_idx: i, + count: samples[i], + mass: mass_of(i, &samples, sample_weights), + }) + .collect() + }) + .collect(); + let mut base_tree = BaseTreeRegressor { nodes, parameters: Some(parameters), @@ -256,8 +293,7 @@ impl, Y: Array1> _phantom_y: PhantomData, }; - let mut visitor = - NodeVisitor::::new(0, samples, sample_weights, &order, x, &y_m, 1); + let mut visitor = NodeVisitor::::new(0, sorted_node_elements, x, &y_m, 1); let mut visitor_queue: LinkedList> = LinkedList::new(); @@ -316,7 +352,7 @@ impl, Y: Array1> ) -> bool { let (_, n_attr) = visitor.x.shape(); - let n: usize = visitor.samples.iter().sum(); + let n: usize = visitor.num_samples(); if n < self.parameters().min_samples_split { return false; @@ -324,6 +360,7 @@ impl, Y: Array1> let sum = self.nodes()[visitor.node].output * mass; + // TODO later: get rid of this allocation in every iteration let mut variables = (0..n_attr).collect::>(); if mtry < n_attr { @@ -360,20 +397,12 @@ impl, Y: Array1> rng: &mut impl rand::Rng, ) { let (min_val, max_val) = { - let mut min_opt = None; - let mut max_opt = None; - for &i in &visitor.order[j] { - if visitor.samples[i] > 0 { - min_opt = Some(*visitor.x.get((i, j))); - break; - } - } - for &i in visitor.order[j].iter().rev() { - if visitor.samples[i] > 0 { - max_opt = Some(*visitor.x.get((i, j))); - break; - } - } + let min_opt = visitor.sorted_node_elements[j] + .first() + .map(|elem| visitor.x.get((elem.row_idx, j))); + let max_opt = visitor.sorted_node_elements[j] + .last() + .map(|elem| visitor.x.get((elem.row_idx, j))); if min_opt.is_none() { return; } @@ -389,15 +418,11 @@ impl, Y: Array1> let mut true_sum = 0f64; let mut true_mass = 0f64; let mut true_count = 0; - for &i in &visitor.order[j] { - if visitor.samples[i] > 0 { - if visitor.x.get((i, j)).to_f64().unwrap() <= split_value { - true_sum += visitor.mass_of(i) * visitor.y.get(i).to_f64().unwrap(); - true_count += visitor.samples[i]; - true_mass += visitor.mass_of(i); - } else { - break; - } + for elem in &visitor.sorted_node_elements[j] { + if visitor.x.get((elem.row_idx, j)).to_f64().unwrap() <= split_value { + true_sum += elem.mass * visitor.y.get(elem.row_idx).to_f64().unwrap(); + true_count += elem.count; + true_mass += elem.mass; } } @@ -448,99 +473,104 @@ impl, Y: Array1> let mut true_mass = 0f64; let mut prevx = Option::None; - for i in visitor.order[j].iter() { - if visitor.samples[*i] > 0 { - let x_ij = *visitor.x.get((*i, j)); + for elem in visitor.sorted_node_elements[j].iter() { + let x_ij = *visitor.x.get((elem.row_idx, j)); - if prevx.is_none() || x_ij == prevx.unwrap() { - prevx = Some(x_ij); - true_count += visitor.samples[*i]; - true_mass += visitor.mass_of(*i); - true_sum += visitor.mass_of(*i) * visitor.y.get(*i).to_f64().unwrap(); - continue; - } + if prevx.is_none() || x_ij == prevx.unwrap() { + prevx = Some(x_ij); + true_count += elem.count; + true_mass += elem.mass; + true_sum += elem.mass * visitor.y.get(elem.row_idx).to_f64().unwrap(); + continue; + } - let false_count = n - true_count; + let false_count = n - true_count; - if true_count < self.parameters().min_samples_leaf - || false_count < self.parameters().min_samples_leaf - { - prevx = Some(x_ij); - true_count += visitor.samples[*i]; - true_mass += visitor.mass_of(*i); - true_sum += visitor.mass_of(*i) * visitor.y.get(*i).to_f64().unwrap(); - continue; - } + if true_count < self.parameters().min_samples_leaf + || false_count < self.parameters().min_samples_leaf + { + prevx = Some(x_ij); + true_count += elem.count; + true_mass += elem.mass; + true_sum += elem.mass * visitor.y.get(elem.row_idx).to_f64().unwrap(); + continue; + } - let true_mean = if true_mass > 0.0 { - true_sum / true_mass - } else { - 0.0 - }; - let false_mass = mass - true_mass; - let false_mean = if false_mass > 0.0 { - (sum - true_sum) / false_mass - } else { - 0.0 - }; - - let gain = (true_mass * true_mean * true_mean - + false_mass * false_mean * false_mean) - - parent_gain; - - if self.nodes()[visitor.node].split_score.is_none() - || gain > self.nodes()[visitor.node].split_score.unwrap() - { - self.nodes[visitor.node].split_feature = j; - self.nodes[visitor.node].split_value = - Option::Some((x_ij + prevx.unwrap()).to_f64().unwrap() / 2f64); - self.nodes[visitor.node].split_score = Option::Some(gain); - - visitor.true_child_output = true_mean; - visitor.false_child_output = false_mean; - } + let true_mean = if true_mass > 0.0 { + true_sum / true_mass + } else { + 0.0 + }; + let false_mass = mass - true_mass; + let false_mean = if false_mass > 0.0 { + (sum - true_sum) / false_mass + } else { + 0.0 + }; - prevx = Some(x_ij); - true_sum += visitor.mass_of(*i) * visitor.y.get(*i).to_f64().unwrap(); - true_count += visitor.samples[*i]; - true_mass += visitor.mass_of(*i); + let gain = (true_mass * true_mean * true_mean + false_mass * false_mean * false_mean) + - parent_gain; + + if self.nodes()[visitor.node].split_score.is_none() + || gain > self.nodes()[visitor.node].split_score.unwrap() + { + self.nodes[visitor.node].split_feature = j; + self.nodes[visitor.node].split_value = + Option::Some((x_ij + prevx.unwrap()).to_f64().unwrap() / 2f64); + self.nodes[visitor.node].split_score = Option::Some(gain); + + visitor.true_child_output = true_mean; + visitor.false_child_output = false_mean; } + + prevx = Some(x_ij); + true_sum += elem.mass * visitor.y.get(elem.row_idx).to_f64().unwrap(); + true_count += elem.count; + true_mass += elem.mass; } } fn split<'a>( &mut self, - mut visitor: NodeVisitor<'a, TX, TY, X, Y>, + visitor: NodeVisitor<'a, TX, TY, X, Y>, mtry: usize, visitor_queue: &mut LinkedList>, rng: &mut impl rand::Rng, ) -> bool { - let (n, _) = visitor.x.shape(); - let mut tc = 0; - let mut fc = 0; - let mut true_mass = 0f64; - let mut false_mass = 0f64; - let mut true_samples: Vec = vec![0; n]; - - for (i, true_sample) in true_samples.iter_mut().enumerate().take(n) { - if visitor.samples[i] > 0 { - if visitor - .x - .get((i, self.nodes()[visitor.node].split_feature)) - .to_f64() - .unwrap() - <= self.nodes()[visitor.node].split_value.unwrap_or(f64::NAN) - { - *true_sample = visitor.samples[i]; - tc += *true_sample; - true_mass += visitor.mass_of(i); - visitor.samples[i] = 0; - } else { - fc += visitor.samples[i]; - false_mass += visitor.mass_of(i); - } - } + let this_node = &self.nodes()[visitor.node]; + + // sorted_node_elements needs to be turned into + // Vec< (Vec, Vec) > and then into + // (Vec, Vec) + + // for each row_index, does it belong in the true branch or not? + let mut is_true = vec![false; visitor.x.shape().0]; + for e in &visitor.sorted_node_elements[0] { + is_true[e.row_idx] = is_true_sample(e, visitor.x, this_node); } + // now use this to partition each of the vectors + let (true_samples, false_samples): (Vec>, Vec>) = visitor + .sorted_node_elements + .iter() + .map(|col| col.iter().partition(|e| is_true[e.row_idx])) + .unzip(); + + let tc: usize = true_samples[0usize] + .iter() + .map(|node_elem| node_elem.count) + .sum::(); + let true_mass = true_samples[0usize] + .iter() + .map(|node_elem| node_elem.mass) + .sum::(); + let fc: usize = false_samples[0usize] + .iter() + .map(|node_elem| node_elem.count) + .sum::(); + let false_mass = false_samples[0usize] + .iter() + .map(|node_elem| node_elem.mass) + .sum::(); if tc < self.parameters().min_samples_leaf || fc < self.parameters().min_samples_leaf { self.nodes[visitor.node].split_feature = 0; @@ -564,8 +594,6 @@ impl, Y: Array1> let mut true_visitor = NodeVisitor::::new( true_child_idx, true_samples, - visitor.sample_weights, - visitor.order, visitor.x, visitor.y, visitor.level + 1, @@ -577,9 +605,7 @@ impl, Y: Array1> let mut false_visitor = NodeVisitor::::new( false_child_idx, - visitor.samples, - visitor.sample_weights, - visitor.order, + false_samples, visitor.x, visitor.y, visitor.level + 1, From 4e4c61c550178b67847fa77ea9f6037120e1f972 Mon Sep 17 00:00:00 2001 From: slievens <5492873+slievens@users.noreply.github.com> Date: Fri, 2 Oct 2026 14:06:45 +0200 Subject: [PATCH 02/17] Less allocation, less sorting --- src/ensemble/base_forest_regressor.rs | 10 ++++++ src/tree/base_tree_regressor.rs | 51 ++++++++++++++++++++------- 2 files changed, 49 insertions(+), 12 deletions(-) diff --git a/src/ensemble/base_forest_regressor.rs b/src/ensemble/base_forest_regressor.rs index 768f64c4..86a21a17 100644 --- a/src/ensemble/base_forest_regressor.rs +++ b/src/ensemble/base_forest_regressor.rs @@ -5,6 +5,7 @@ use std::fmt::Debug; use serde::{Deserialize, Serialize}; use crate::error::{Failed, FailedError}; +use crate::linalg::basic::arrays::MutArrayView1; use crate::linalg::basic::arrays::{Array1, Array2}; use crate::numbers::basenum::Number; use crate::numbers::floatnum::FloatNumber; @@ -120,6 +121,14 @@ impl, Y: Array1 }) .transpose()?; + // Compute the order of each attribute once + let mut order: Vec> = Vec::with_capacity(num_attributes); + + for i in 0..num_attributes { + let mut col_i: Vec = x.get_col(i).iterator(0).copied().collect(); + order.push(col_i.argsort_mut()); + } + for _ in 0..parameters.n_trees { if parameters.bootstrap { samples = BaseForestRegressor::::sample_with_replacement( @@ -147,6 +156,7 @@ impl, Y: Array1 sample_weights, samples.clone(), mtry, + &order, params, )?; trees.push(tree); diff --git a/src/tree/base_tree_regressor.rs b/src/tree/base_tree_regressor.rs index 0fe54bae..069fb04f 100644 --- a/src/tree/base_tree_regressor.rs +++ b/src/tree/base_tree_regressor.rs @@ -223,6 +223,14 @@ impl, Y: Array1> )); } + // Compute the order of each attribute once + let mut order: Vec> = Vec::new(); + + for i in 0..num_attributes { + let mut col_i: Vec = x.get_col(i).iterator(0).copied().collect(); + order.push(col_i.argsort_mut()); + } + let samples = vec![1; x_nrows]; BaseTreeRegressor::fit_weak_learner( x, @@ -230,6 +238,7 @@ impl, Y: Array1> sample_weights, samples, num_attributes, + &order, parameters, ) } @@ -240,12 +249,12 @@ impl, Y: Array1> sample_weights: Option<&[f64]>, samples: Vec, mtry: usize, + order: &[Vec], parameters: BaseTreeRegressorParameters, ) -> Result, Failed> { let y_m = y.clone(); let y_ncols = y_m.shape(); - let (_, num_attributes) = x.shape(); let mut nodes: Vec = Vec::new(); let mut rng = get_rng_impl(parameters.seed); @@ -261,12 +270,6 @@ impl, Y: Array1> let root = Node::new(sum / mass); nodes.push(root); - let mut order: Vec> = Vec::new(); - - for i in 0..num_attributes { - let mut col_i: Vec = x.get_col(i).iterator(0).copied().collect(); - order.push(col_i.argsort_mut()); - } let sorted_node_elements: Vec> = order .iter() @@ -301,10 +304,17 @@ impl, Y: Array1> visitor_queue.push_back(visitor); } + let mut scratch_buffer = vec![false; x.shape().0]; let max_depth = base_tree.parameters().max_depth.unwrap_or(u16::MAX); while let Some(node) = visitor_queue.pop_front() { if node.level < max_depth { - base_tree.split(node, mtry, &mut visitor_queue, &mut rng); + base_tree.split( + node, + mtry, + &mut visitor_queue, + &mut rng, + &mut scratch_buffer, + ); } } @@ -536,6 +546,7 @@ impl, Y: Array1> mtry: usize, visitor_queue: &mut LinkedList>, rng: &mut impl rand::Rng, + buffer: &mut [bool], // buffer used to track the splitting ) -> bool { let this_node = &self.nodes()[visitor.node]; @@ -544,15 +555,31 @@ impl, Y: Array1> // (Vec, Vec) // for each row_index, does it belong in the true branch or not? - let mut is_true = vec![false; visitor.x.shape().0]; + let is_true = buffer; + let mut n_true = 0usize; for e in &visitor.sorted_node_elements[0] { - is_true[e.row_idx] = is_true_sample(e, visitor.x, this_node); + let t = is_true_sample(e, visitor.x, this_node); + is_true[e.row_idx] = t; + n_true += t as usize; } - // now use this to partition each of the vectors + let n_false = visitor.sorted_node_elements[0].len() - n_true; + + // now use this to partition each of the vectors. Preallocate vectors to avoid reallocations let (true_samples, false_samples): (Vec>, Vec>) = visitor .sorted_node_elements .iter() - .map(|col| col.iter().partition(|e| is_true[e.row_idx])) + .map(|col| { + let mut true_vec = Vec::with_capacity(n_true); // preallocate + let mut false_vec = Vec::with_capacity(n_false); + for e in col { + if is_true[e.row_idx] { + true_vec.push(*e); + } else { + false_vec.push(*e); + } + } + (true_vec, false_vec) + }) .unzip(); let tc: usize = true_samples[0usize] From 2a427b3818594175f41b3cb0ecfe2b663db98225 Mon Sep 17 00:00:00 2001 From: slievens <5492873+slievens@users.noreply.github.com> Date: Fri, 2 Oct 2026 15:57:01 +0200 Subject: [PATCH 03/17] Node as (start, end) --- src/tree/base_tree_regressor.rs | 236 ++++++++++++++++++++------------ 1 file changed, 149 insertions(+), 87 deletions(-) diff --git a/src/tree/base_tree_regressor.rs b/src/tree/base_tree_regressor.rs index 069fb04f..2ccc4963 100644 --- a/src/tree/base_tree_regressor.rs +++ b/src/tree/base_tree_regressor.rs @@ -131,7 +131,9 @@ struct NodeVisitor<'a, TX: Number + PartialOrd, TY: Number, X: Array2, Y: Ar y: &'a Y, node: usize, // holds the elements for this node, sorted for each feature [num_features, num_samples] - sorted_node_elements: Vec>, + //sorted_node_elements: Vec>, + start_idx: usize, + end_idx: usize, true_child_output: f64, false_child_output: f64, level: u16, @@ -144,7 +146,8 @@ impl<'a, TX: Number + PartialOrd, TY: Number, X: Array2, Y: Array1> { fn new( node_id: usize, - sorted_node_elements: Vec>, + start_idx: usize, + end_idx: usize, x: &'a X, y: &'a Y, level: u16, @@ -153,7 +156,8 @@ impl<'a, TX: Number + PartialOrd, TY: Number, X: Array2, Y: Array1> x, y, node: node_id, - sorted_node_elements, + start_idx, + end_idx, true_child_output: 0f64, false_child_output: 0f64, level, @@ -161,14 +165,6 @@ impl<'a, TX: Number + PartialOrd, TY: Number, X: Array2, Y: Array1> _phantom_ty: PhantomData, } } - - /// number of samples in this node (nodevisitor) - fn num_samples(&self) -> usize { - self.sorted_node_elements[0] - .iter() - .map(|node_elem| node_elem.count) - .sum() - } } /// Weighted count of sample `i`. The weight is 1.0 if no weights are given. @@ -180,29 +176,83 @@ fn mass_of(i: usize, samples: &[usize], sample_weights: Option<&[f64]>) -> f64 { } // Struct representing an element that belongs logically to a Node, as stored in NodeVisitor -#[derive(Copy, Clone)] +#[derive(Copy, Clone, Default)] struct NodeElement { // the row index in the dataset - pub row_idx: usize, + pub row_idx: u32, // the number of times this row is present, should always be > 0 - pub count: usize, + pub count: u32, // total mass of this element, equals count * mass of individual element. // equals count when no sample weights were used pub mass: f64, } +impl NodeElement { + // return the row_idx as usize + #[inline(always)] + fn row(&self) -> usize { + self.row_idx as usize + } +} + // Checks whether the example indicated by node_element is a "true child" for this node fn is_true_sample(node_element: &NodeElement, x: &X, node: &Node) -> bool where TX: Number + PartialOrd, X: Array2, { - x.get((node_element.row_idx, node.split_feature)) + x.get((node_element.row(), node.split_feature)) .to_f64() .unwrap() <= node.split_value.unwrap_or(f64::NAN) } +// slice: slice that will be partitioned +// scratch: temp buffer +// is_true: is_true[idx] checks whether element with row idx equal to idx belongs to the true branch +// returns: index of first element of false branch +fn stable_partition( + slice: &mut [NodeElement], + scratch: &mut [NodeElement], + is_true: &[bool], +) -> usize { + // Note: this is intentionally written without an if/else branch in the main loop + let n = slice.len(); + let scratch = &mut scratch[..n]; + let (mut w, mut f) = (0usize, 0usize); + for i in 0..n { + let e = slice[i]; + let t = is_true[e.row_idx as usize]; + slice[w] = e; // w <= i, so this never clobbers an unread element + scratch[f] = e; + // advance only one of the pointers + w += t as usize; + f += (!t) as usize; + } + slice[w..].copy_from_slice(&scratch[..f]); + w +} + +struct ScratchPad { + is_true: Vec, + node_elements: Vec, + shared_node_elements: Vec>, +} + +impl ScratchPad { + fn new( + is_true: Vec, + node_elements: Vec, + shared_node_elements: Vec>, + ) -> Self { + Self { + is_true, + node_elements, + shared_node_elements, + } + } +} + impl, Y: Array1> BaseTreeRegressor { @@ -271,20 +321,27 @@ impl, Y: Array1> let root = Node::new(sum / mass); nodes.push(root); - let sorted_node_elements: Vec> = order + let shared_node_elements: Vec> = order .iter() .map(|col_order| { col_order .iter() .filter(|&&i| samples[i] > 0) .map(|&i| NodeElement { - row_idx: i, - count: samples[i], + row_idx: i as u32, + count: samples[i] as u32, mass: mass_of(i, &samples, sample_weights), }) .collect() }) .collect(); + let end_idx = shared_node_elements[0].len(); + + let mut scratch_pad = ScratchPad::new( + vec![false; x.shape().0], + vec![NodeElement::default(); end_idx], + shared_node_elements, + ); let mut base_tree = BaseTreeRegressor { nodes, @@ -296,25 +353,18 @@ impl, Y: Array1> _phantom_y: PhantomData, }; - let mut visitor = NodeVisitor::::new(0, sorted_node_elements, x, &y_m, 1); + let mut visitor = NodeVisitor::::new(0, 0, end_idx, x, &y_m, 1); let mut visitor_queue: LinkedList> = LinkedList::new(); - if base_tree.find_best_cutoff(&mut visitor, mtry, mass, &mut rng) { + if base_tree.find_best_cutoff(&mut visitor, mtry, mass, &mut rng, &scratch_pad) { visitor_queue.push_back(visitor); } - let mut scratch_buffer = vec![false; x.shape().0]; let max_depth = base_tree.parameters().max_depth.unwrap_or(u16::MAX); while let Some(node) = visitor_queue.pop_front() { if node.level < max_depth { - base_tree.split( - node, - mtry, - &mut visitor_queue, - &mut rng, - &mut scratch_buffer, - ); + base_tree.split(node, mtry, &mut visitor_queue, &mut rng, &mut scratch_pad); } } @@ -359,10 +409,15 @@ impl, Y: Array1> mtry: usize, mass: f64, rng: &mut impl rand::Rng, + scratch_pad: &ScratchPad, ) -> bool { let (_, n_attr) = visitor.x.shape(); - let n: usize = visitor.num_samples(); + //let n: usize = visitor.num_samples(); + let n: usize = scratch_pad.shared_node_elements[0][visitor.start_idx..visitor.end_idx] + .iter() + .map(|elem| elem.count as usize) + .sum(); if n < self.parameters().min_samples_split { return false; @@ -385,10 +440,27 @@ impl, Y: Array1> for variable in variables.iter().take(mtry) { match splitter { Splitter::Random => { - self.find_random_split(visitor, n, mass, sum, parent_gain, *variable, rng); + self.find_random_split( + visitor, + n, + mass, + sum, + parent_gain, + *variable, + rng, + scratch_pad, + ); } Splitter::Best => { - self.find_best_split(visitor, n, mass, sum, parent_gain, *variable); + self.find_best_split( + visitor, + n, + mass, + sum, + parent_gain, + *variable, + scratch_pad, + ); } } } @@ -405,19 +477,15 @@ impl, Y: Array1> parent_gain: f64, j: usize, rng: &mut impl rand::Rng, + scratch_pad: &ScratchPad, ) { - let (min_val, max_val) = { - let min_opt = visitor.sorted_node_elements[j] - .first() - .map(|elem| visitor.x.get((elem.row_idx, j))); - let max_opt = visitor.sorted_node_elements[j] - .last() - .map(|elem| visitor.x.get((elem.row_idx, j))); - if min_opt.is_none() { - return; - } - (min_opt.unwrap(), max_opt.unwrap()) - }; + if visitor.start_idx == visitor.end_idx { + return; + } + let first_elem = scratch_pad.shared_node_elements[j][visitor.start_idx]; + let min_val = visitor.x.get((first_elem.row(), j)); + let last_elem = scratch_pad.shared_node_elements[j][visitor.end_idx - 1]; + let max_val = visitor.x.get((last_elem.row(), j)); if min_val >= max_val { return; @@ -428,17 +496,17 @@ impl, Y: Array1> let mut true_sum = 0f64; let mut true_mass = 0f64; let mut true_count = 0; - for elem in &visitor.sorted_node_elements[j] { - if visitor.x.get((elem.row_idx, j)).to_f64().unwrap() <= split_value { - true_sum += elem.mass * visitor.y.get(elem.row_idx).to_f64().unwrap(); + for elem in &scratch_pad.shared_node_elements[j][visitor.start_idx..visitor.end_idx] { + if visitor.x.get((elem.row(), j)).to_f64().unwrap() <= split_value { + true_sum += elem.mass * visitor.y.get(elem.row()).to_f64().unwrap(); true_count += elem.count; true_mass += elem.mass; } } - let false_count = n - true_count; + let false_count = n - (true_count as usize); - if true_count < self.parameters().min_samples_leaf + if (true_count as usize) < self.parameters().min_samples_leaf || false_count < self.parameters().min_samples_leaf { return; @@ -477,32 +545,33 @@ impl, Y: Array1> sum: f64, parent_gain: f64, j: usize, + scratch_pad: &ScratchPad, ) { let mut true_sum = 0f64; let mut true_count = 0; let mut true_mass = 0f64; let mut prevx = Option::None; - for elem in visitor.sorted_node_elements[j].iter() { - let x_ij = *visitor.x.get((elem.row_idx, j)); + for elem in &scratch_pad.shared_node_elements[j][visitor.start_idx..visitor.end_idx] { + let x_ij = *visitor.x.get((elem.row(), j)); if prevx.is_none() || x_ij == prevx.unwrap() { prevx = Some(x_ij); true_count += elem.count; true_mass += elem.mass; - true_sum += elem.mass * visitor.y.get(elem.row_idx).to_f64().unwrap(); + true_sum += elem.mass * visitor.y.get(elem.row()).to_f64().unwrap(); continue; } - let false_count = n - true_count; + let false_count = n - (true_count as usize); - if true_count < self.parameters().min_samples_leaf + if (true_count as usize) < self.parameters().min_samples_leaf || false_count < self.parameters().min_samples_leaf { prevx = Some(x_ij); true_count += elem.count; true_mass += elem.mass; - true_sum += elem.mass * visitor.y.get(elem.row_idx).to_f64().unwrap(); + true_sum += elem.mass * visitor.y.get(elem.row()).to_f64().unwrap(); continue; } @@ -534,7 +603,7 @@ impl, Y: Array1> } prevx = Some(x_ij); - true_sum += elem.mass * visitor.y.get(elem.row_idx).to_f64().unwrap(); + true_sum += elem.mass * visitor.y.get(elem.row()).to_f64().unwrap(); true_count += elem.count; true_mass += elem.mass; } @@ -546,7 +615,7 @@ impl, Y: Array1> mtry: usize, visitor_queue: &mut LinkedList>, rng: &mut impl rand::Rng, - buffer: &mut [bool], // buffer used to track the splitting + scratch_pad: &mut ScratchPad, // buffer used to track the splitting ) -> bool { let this_node = &self.nodes()[visitor.node]; @@ -555,46 +624,37 @@ impl, Y: Array1> // (Vec, Vec) // for each row_index, does it belong in the true branch or not? - let is_true = buffer; + let is_true = &mut scratch_pad.is_true; let mut n_true = 0usize; - for e in &visitor.sorted_node_elements[0] { + for e in &scratch_pad.shared_node_elements[0][visitor.start_idx..visitor.end_idx] { let t = is_true_sample(e, visitor.x, this_node); - is_true[e.row_idx] = t; + is_true[e.row()] = t; n_true += t as usize; } - let n_false = visitor.sorted_node_elements[0].len() - n_true; - // now use this to partition each of the vectors. Preallocate vectors to avoid reallocations - let (true_samples, false_samples): (Vec>, Vec>) = visitor - .sorted_node_elements - .iter() - .map(|col| { - let mut true_vec = Vec::with_capacity(n_true); // preallocate - let mut false_vec = Vec::with_capacity(n_false); - for e in col { - if is_true[e.row_idx] { - true_vec.push(*e); - } else { - false_vec.push(*e); - } - } - (true_vec, false_vec) - }) - .unzip(); + for j in 0..visitor.x.shape().1 { + stable_partition( + &mut scratch_pad.shared_node_elements[j][visitor.start_idx..visitor.end_idx], + &mut scratch_pad.node_elements, + is_true, + ); + } + + let split_idx = visitor.start_idx + n_true; - let tc: usize = true_samples[0usize] + let tc: usize = scratch_pad.shared_node_elements[0][visitor.start_idx..split_idx] .iter() - .map(|node_elem| node_elem.count) + .map(|node_elem| node_elem.count as usize) .sum::(); - let true_mass = true_samples[0usize] + let true_mass = scratch_pad.shared_node_elements[0][visitor.start_idx..split_idx] .iter() .map(|node_elem| node_elem.mass) .sum::(); - let fc: usize = false_samples[0usize] + let fc: usize = scratch_pad.shared_node_elements[0][split_idx..visitor.end_idx] .iter() - .map(|node_elem| node_elem.count) + .map(|node_elem| node_elem.count as usize) .sum::(); - let false_mass = false_samples[0usize] + let false_mass = scratch_pad.shared_node_elements[0][split_idx..visitor.end_idx] .iter() .map(|node_elem| node_elem.mass) .sum::(); @@ -620,25 +680,27 @@ impl, Y: Array1> let mut true_visitor = NodeVisitor::::new( true_child_idx, - true_samples, + visitor.start_idx, + split_idx, visitor.x, visitor.y, visitor.level + 1, ); - if self.find_best_cutoff(&mut true_visitor, mtry, true_mass, rng) { + if self.find_best_cutoff(&mut true_visitor, mtry, true_mass, rng, scratch_pad) { visitor_queue.push_back(true_visitor); } let mut false_visitor = NodeVisitor::::new( false_child_idx, - false_samples, + split_idx, + visitor.end_idx, visitor.x, visitor.y, visitor.level + 1, ); - if self.find_best_cutoff(&mut false_visitor, mtry, false_mass, rng) { + if self.find_best_cutoff(&mut false_visitor, mtry, false_mass, rng, scratch_pad) { visitor_queue.push_back(false_visitor); } From 143edc8f83182596fbf64361a5103bae18150f69 Mon Sep 17 00:00:00 2001 From: slievens <5492873+slievens@users.noreply.github.com> Date: Fri, 2 Oct 2026 16:46:41 +0200 Subject: [PATCH 04/17] Avoid partitioning when child nodes can not be split --- src/tree/base_tree_regressor.rs | 77 ++++++++++++++++++--------------- 1 file changed, 41 insertions(+), 36 deletions(-) diff --git a/src/tree/base_tree_regressor.rs b/src/tree/base_tree_regressor.rs index 2ccc4963..ef1dd372 100644 --- a/src/tree/base_tree_regressor.rs +++ b/src/tree/base_tree_regressor.rs @@ -1,4 +1,4 @@ -use std::collections::LinkedList; +use std::collections::VecDeque; use std::default::Default; use std::fmt::Debug; use std::marker::PhantomData; @@ -302,9 +302,7 @@ impl, Y: Array1> order: &[Vec], parameters: BaseTreeRegressorParameters, ) -> Result, Failed> { - let y_m = y.clone(); - - let y_ncols = y_m.shape(); + let n_rows = y.shape(); let mut nodes: Vec = Vec::new(); let mut rng = get_rng_impl(parameters.seed); @@ -312,10 +310,10 @@ impl, Y: Array1> let mut sum = 0f64; let mut mass = 0f64; - for i in 0..y_ncols { + for i in 0..n_rows { let mass_i = mass_of(i, &samples, sample_weights); mass += mass_i; - sum += mass_i * y_m.get(i).to_f64().unwrap(); + sum += mass_i * y.get(i).to_f64().unwrap(); } let root = Node::new(sum / mass); @@ -353,9 +351,9 @@ impl, Y: Array1> _phantom_y: PhantomData, }; - let mut visitor = NodeVisitor::::new(0, 0, end_idx, x, &y_m, 1); + let mut visitor = NodeVisitor::::new(0, 0, end_idx, x, y, 1); - let mut visitor_queue: LinkedList> = LinkedList::new(); + let mut visitor_queue: VecDeque> = VecDeque::new(); if base_tree.find_best_cutoff(&mut visitor, mtry, mass, &mut rng, &scratch_pad) { visitor_queue.push_back(visitor); @@ -613,7 +611,7 @@ impl, Y: Array1> &mut self, visitor: NodeVisitor<'a, TX, TY, X, Y>, mtry: usize, - visitor_queue: &mut LinkedList>, + visitor_queue: &mut VecDeque>, rng: &mut impl rand::Rng, scratch_pad: &mut ScratchPad, // buffer used to track the splitting ) -> bool { @@ -623,6 +621,10 @@ impl, Y: Array1> // Vec< (Vec, Vec) > and then into // (Vec, Vec) + let mut tc = 0usize; + let mut true_mass = 0f64; + let mut fc = 0usize; + let mut false_mass = 0f64; // for each row_index, does it belong in the true branch or not? let is_true = &mut scratch_pad.is_true; let mut n_true = 0usize; @@ -630,35 +632,17 @@ impl, Y: Array1> let t = is_true_sample(e, visitor.x, this_node); is_true[e.row()] = t; n_true += t as usize; + // Fill in tc, etc while we are at it + if t { + tc += e.count as usize; + true_mass += e.mass; + } else { + fc += e.count as usize; + false_mass += e.mass; + } } - for j in 0..visitor.x.shape().1 { - stable_partition( - &mut scratch_pad.shared_node_elements[j][visitor.start_idx..visitor.end_idx], - &mut scratch_pad.node_elements, - is_true, - ); - } - - let split_idx = visitor.start_idx + n_true; - - let tc: usize = scratch_pad.shared_node_elements[0][visitor.start_idx..split_idx] - .iter() - .map(|node_elem| node_elem.count as usize) - .sum::(); - let true_mass = scratch_pad.shared_node_elements[0][visitor.start_idx..split_idx] - .iter() - .map(|node_elem| node_elem.mass) - .sum::(); - let fc: usize = scratch_pad.shared_node_elements[0][split_idx..visitor.end_idx] - .iter() - .map(|node_elem| node_elem.count as usize) - .sum::(); - let false_mass = scratch_pad.shared_node_elements[0][split_idx..visitor.end_idx] - .iter() - .map(|node_elem| node_elem.mass) - .sum::(); - + // Stop early if it is clear that there will be too few examples in the leaf if tc < self.parameters().min_samples_leaf || fc < self.parameters().min_samples_leaf { self.nodes[visitor.node].split_feature = 0; self.nodes[visitor.node].split_value = Option::None; @@ -667,6 +651,9 @@ impl, Y: Array1> return false; } + // Add the child nodes to the tree + let split_idx = visitor.start_idx + n_true; + let true_child_idx = self.nodes().len(); self.nodes.push(Node::new(visitor.true_child_output)); @@ -678,6 +665,24 @@ impl, Y: Array1> self.depth = u16::max(self.depth, visitor.level + 1); + // If the child nodes can not be split any further, there is no point is partitioning the ranges + let max_depth = self.parameters().max_depth.unwrap_or(u16::MAX); + let child_level = visitor.level + 1; + let min_split = self.parameters().min_samples_split; + let true_can_split = child_level < max_depth && tc >= min_split; + let false_can_split = child_level < max_depth && fc >= min_split; + if !true_can_split && !false_can_split { + return true; // both children are leaves: no partition, no search + } + + for j in 0..visitor.x.shape().1 { + stable_partition( + &mut scratch_pad.shared_node_elements[j][visitor.start_idx..visitor.end_idx], + &mut scratch_pad.node_elements, + is_true, + ); + } + let mut true_visitor = NodeVisitor::::new( true_child_idx, visitor.start_idx, From 70296be892939a0f5f5c3657ede1d4ac476810c4 Mon Sep 17 00:00:00 2001 From: slievens <5492873+slievens@users.noreply.github.com> Date: Fri, 2 Oct 2026 17:49:19 +0200 Subject: [PATCH 05/17] try to avoid numerical errors --- src/tree/base_tree_regressor.rs | 5 +++-- 1 file changed, 3 insertions(+), 2 deletions(-) diff --git a/src/tree/base_tree_regressor.rs b/src/tree/base_tree_regressor.rs index ef1dd372..1cb7c46f 100644 --- a/src/tree/base_tree_regressor.rs +++ b/src/tree/base_tree_regressor.rs @@ -592,8 +592,9 @@ impl, Y: Array1> || gain > self.nodes()[visitor.node].split_score.unwrap() { self.nodes[visitor.node].split_feature = j; - self.nodes[visitor.node].split_value = - Option::Some((x_ij + prevx.unwrap()).to_f64().unwrap() / 2f64); + self.nodes[visitor.node].split_value = Option::Some( + (x_ij.to_f64().unwrap() + prevx.unwrap().to_f64().unwrap()) / 2f64, + ); self.nodes[visitor.node].split_score = Option::Some(gain); visitor.true_child_output = true_mean; From 91145787531e2d1f79e01a3eb3a98f66bd53cf30 Mon Sep 17 00:00:00 2001 From: slievens <5492873+slievens@users.noreply.github.com> Date: Fri, 2 Oct 2026 21:13:16 +0200 Subject: [PATCH 06/17] Add tests for sklearn parity of randomforestregressor. Added a test (with random data) to see whether the implemenation matches that of sklearn. Discovered an existing bug in the implementation that made that all trees of the forest used the same "random" features. An additional test for this was also added: it tests that without bootstrapping and with a single feature, the trees are still different. This test failed with the "old" implementation. --- src/ensemble/base_forest_regressor.rs | 39 ++++- src/ensemble/random_forest_classifier.rs | 2 +- src/ensemble/random_forest_regressor.rs | 192 +++++++++++++++++++++++ 3 files changed, 231 insertions(+), 2 deletions(-) diff --git a/src/ensemble/base_forest_regressor.rs b/src/ensemble/base_forest_regressor.rs index 86a21a17..e91ead3e 100644 --- a/src/ensemble/base_forest_regressor.rs +++ b/src/ensemble/base_forest_regressor.rs @@ -147,7 +147,7 @@ impl, Y: Array1 max_depth: parameters.max_depth, min_samples_leaf: parameters.min_samples_leaf, min_samples_split: parameters.min_samples_split, - seed: Some(parameters.seed), + seed: Some(rng.random::()), // give each tree a different seed for its rng splitter: parameters.splitter.clone(), }; let tree = BaseTreeRegressor::fit_weak_learner( @@ -439,6 +439,43 @@ mod tests { } } + #[test] + fn each_tree_gets_different_feature_sample() { + // Without bootstrap, all trees get the same rows. With m = 1, each node uses one + // random feature, thus the trees must be different. If all trees use the same seed, + // all trees are equal. The Random splitter also uses this rng for the thresholds. + + // Create some data + let n_rows = 30; + let x: DenseMatrix = DenseMatrix::from_iterator( + (0..4 * n_rows).map(|k| ((k * 7919) % 101) as f64), + n_rows, + 4, + 0, + ); + let y: Vec = (0..n_rows).map(|i| ((i * 31) % 17) as f64).collect(); + + for splitter in [Splitter::Best, Splitter::Random] { + let params = BaseForestRegressorParameters { + max_depth: None, + min_samples_leaf: 1, + min_samples_split: 2, + n_trees: 10, + m: Some(1), + keep_samples: false, + seed: 42, + bootstrap: false, + splitter: splitter.clone(), + }; + let forest = BaseForestRegressor::fit(&x, &y, None, params).unwrap(); + let trees = forest.trees.unwrap(); + assert!( + trees.iter().any(|tree| tree != &trees[0]), + "all trees are equal (splitter: {splitter:?})" + ); + } + } + #[test] fn fit_with_weights_predicts_approx_weighted_mean() { // 20 rows, 1 feature. Stumps (max_depth = 0): each tree predicts the diff --git a/src/ensemble/random_forest_classifier.rs b/src/ensemble/random_forest_classifier.rs index ff7a4329..03512155 100644 --- a/src/ensemble/random_forest_classifier.rs +++ b/src/ensemble/random_forest_classifier.rs @@ -503,7 +503,7 @@ impl, Y: Array1()), // give each tree a different seed for its rng }; let tree = DecisionTreeClassifier::fit_weak_learner(x, y, samples, mtry, params)?; trees.push(tree); diff --git a/src/ensemble/random_forest_regressor.rs b/src/ensemble/random_forest_regressor.rs index 95eeb483..9ebcc53c 100644 --- a/src/ensemble/random_forest_regressor.rs +++ b/src/ensemble/random_forest_regressor.rs @@ -797,4 +797,196 @@ mod tests { let msg = "'fit' should be called before calling 'predict'"; assert_eq!(yhat.err(), Some(Failed::predict(msg))); } + + mod sklearn_parity { + use super::*; + // sklearn parity tests. + // + // smartcore and numpy use different RNGs, thus the bootstrap samples and the feature + // subsets are different. With many trees, both forests converge to the same bagged + // predictor. Thus we compare the predictions within a tolerance. + // + // Reference: sklearn 1.9.1, numpy 2.4.6. Each reference value is the mean of 10 sklearn + // runs (random_state = 0..10). The tolerance is approximately 4 x the largest standard + // deviation of one sklearn run, for each row. + // + // ```python + // rng = np.random.default_rng(0) + // x = np.round(rng.uniform(-1, 1, (40, 4)), 4) + // y = np.round(x[:, 0] * x[:, 1] + np.sin(3 * x[:, 2]) + 0.1 * rng.normal(size=40), 4) + // x_probe = np.round(rng.uniform(-1, 1, (10, 4)), 4) + // for max_features in [1.0, 2]: + // for seed in range(10): + // rf = RandomForestRegressor(n_estimators=2000, max_features=max_features, + // max_depth=None, min_samples_leaf=1, min_samples_split=2, bootstrap=True, + // oob_score=True, random_state=seed).fit(x, y) + // # collect rf.predict(x), rf.predict(x_probe), rf.oob_prediction_ + // ``` + // + // TODO: no weighted case. smartcore uses the sample weights for the bootstrap and also + // multiplies them into the tree mass. sklearn uses them only for the bootstrap. + + fn sklearn_parity_train_data() -> (DenseMatrix, Vec) { + let x = DenseMatrix::from_2d_array(&[ + &[0.2739, -0.4604, -0.9181, -0.9669], + &[0.6265, 0.8255, 0.2133, 0.459], + &[0.0872, 0.8701, 0.6317, -0.9945], + &[0.7148, -0.9328, 0.4593, -0.6487], + &[0.7264, 0.0829, -0.4006, -0.1546], + &[-0.9434, -0.7514, 0.3412, 0.2944], + &[0.2308, -0.2326, 0.9944, 0.9617], + &[0.3711, 0.3009, 0.3769, -0.2222], + &[-0.7298, 0.443, 0.0507, -0.3795], + &[-0.0283, 0.779, 0.8681, -0.2844], + &[0.1431, -0.3563, 0.1886, -0.3242], + &[-0.2168, 0.7805, -0.5457, 0.2464], + &[-0.832, 0.6653, 0.5742, -0.5213], + &[0.753, -0.8829, -0.3278, -0.6994], + &[-0.0993, 0.5926, -0.5387, -0.896], + &[-0.1909, -0.603, -0.8185, 0.1607], + &[-0.4026, 0.344, -0.601, 0.8842], + &[-0.2698, -0.789, 0.2582, 0.8543], + &[-0.1192, 0.9092, -0.0002, -0.1495], + &[0.2404, 0.9902, 0.8979, -0.0799], + &[0.5155, -0.0052, 0.0586, 0.5716], + &[-0.1707, 0.469, 0.4223, 0.8641], + &[-0.7701, 0.458, 0.8548, 0.9359], + &[-0.9706, 0.7273, 0.9624, 0.9144], + &[-0.7025, 0.9453, 0.7799, 0.6447], + &[-0.04, -0.5353, 0.6038, 0.8471], + &[-0.4677, 0.0779, -0.1145, 0.862], + &[-0.919, 0.464, 0.2287, -0.9433], + &[0.4384, -0.968, 0.5159, 0.0255], + &[0.8582, -0.8678, 0.6826, -0.8666], + &[-0.3114, -0.1394, 0.9321, 0.1245], + &[-0.4823, -0.5166, 0.7762, -0.5483], + &[-0.7509, -0.4233, 0.1722, 0.1082], + &[0.6194, 0.121, -0.4232, -0.1742], + &[0.6362, 0.253, 0.9182, -0.2612], + &[0.1052, 0.1878, 0.6966, -0.7091], + &[-0.187, 0.8199, -0.9139, 0.6454], + &[-0.1692, 0.6596, -0.9801, -0.2699], + &[-0.8427, 0.3052, -0.4523, 0.4053], + &[0.8876, -0.7464, 0.7296, -0.8811], + ]) + .unwrap(); + let y = vec![ + -0.5144, 0.9957, 0.7839, 0.3660, -0.9022, 1.5099, 0.0804, 1.1980, -0.1768, 0.4984, + 0.3364, -1.0023, 0.5267, -1.3905, -1.0531, -0.4267, -1.0746, 0.9736, -0.1242, + 0.5237, 0.2751, 0.6806, 0.1690, -0.4747, -0.0497, 1.0539, -0.3933, 0.1634, 0.6273, + 0.0960, 0.5208, 1.0106, 0.7643, -1.0745, 0.4076, 0.9968, -0.5477, -0.3399, -1.0701, + 0.0243, + ]; + (x, y) + } + + fn sklearn_parity_probe_data() -> DenseMatrix { + DenseMatrix::from_2d_array(&[ + &[-0.3606, -0.625, 0.3451, -0.6098], + &[0.1554, 0.2045, 0.9248, -0.8555], + &[-0.0001, 0.4882, -0.6455, -0.2239], + &[-0.8742, 0.4518, -0.8245, -0.2098], + &[0.747, -0.0554, 0.8252, 0.5318], + &[0.8306, -0.7452, -0.8529, -0.8593], + &[0.7377, 0.2681, -0.0069, -0.6729], + &[0.3475, -0.364, 0.4218, -0.0793], + &[0.0149, 0.5793, -0.8145, 0.1575], + &[-0.6055, 0.6163, -0.0223, 0.9774], + ]) + .unwrap() + } + + fn assert_close_to_sklearn(actual: &[f64], expected: &[f64], tol: f64, label: &str) { + assert_eq!(actual.len(), expected.len(), "{label}: length"); + for (i, (a, e)) in actual.iter().zip(expected.iter()).enumerate() { + assert!( + (a - e).abs() <= tol, + "{label}, row {i}: smartcore {a}, sklearn {e}, tol {tol}" + ); + } + } + + /// Fits the forest with 2000 trees and compares the predictions with sklearn. + fn check_sklearn_parity( + m: usize, + train_ref: &[f64], + probe_ref: &[f64], + oob_ref: &[f64], + tol: f64, + oob_tol: f64, + ) { + let (x, y) = sklearn_parity_train_data(); + let x_probe = sklearn_parity_probe_data(); + + let parameters = RandomForestRegressorParameters::default() + .with_n_trees(2000) + .with_m(m) + .with_min_samples_leaf(1) + .with_min_samples_split(2) + .with_keep_samples(true) + .with_seed(42); + let forest = RandomForestRegressor::fit(&x, &y, parameters).unwrap(); + + let y_hat: Vec = forest.predict(&x).unwrap(); + assert_close_to_sklearn(&y_hat, train_ref, tol, "train"); + + let y_hat_probe: Vec = forest.predict(&x_probe).unwrap(); + assert_close_to_sklearn(&y_hat_probe, probe_ref, tol, "probe"); + + let y_hat_oob: Vec = forest.predict_oob(&x).unwrap(); + assert_close_to_sklearn(&y_hat_oob, oob_ref, oob_tol, "oob"); + } + + #[test] + fn sklearn_parity_all_features() { + // sklearn max_features = 1.0. Largest std of one run: train 0.0132, probe 0.0097, + // oob 0.0290. + let train_ref = [ + -0.563215, 0.790441, 0.671600, 0.447413, -0.962337, 1.182714, 0.146502, 0.956499, + -0.019351, 0.475851, 0.536221, -0.921016, 0.566884, -1.182586, -0.953829, + -0.492116, -0.970510, 0.906254, -0.122531, 0.492432, 0.111251, 0.701316, 0.166136, + -0.090198, 0.168000, 0.917484, -0.318774, 0.452445, 0.701507, 0.213636, 0.519599, + 0.777545, 0.710208, -1.026581, 0.428588, 0.814387, -0.550148, -0.449774, -0.966291, + 0.151257, + ]; + let probe_ref = [ + 0.840299, 0.478916, -0.941862, -0.540659, 0.284103, -0.685080, -0.272447, 0.791806, + -0.562334, -0.195136, + ]; + let oob_ref = [ + -0.649343, 0.431595, 0.475531, 0.590899, -1.064646, 0.606166, 0.261989, 0.535072, + 0.255804, 0.436133, 0.880022, -0.782491, 0.635705, -0.820807, -0.774075, -0.607414, + -0.781085, 0.788450, -0.119615, 0.437044, -0.173755, 0.739059, 0.160968, 0.594925, + 0.542656, 0.672530, -0.189747, 0.961226, 0.830511, 0.416988, 0.517550, 0.376026, + 0.616485, -0.942362, 0.465415, 0.488966, -0.554237, -0.643990, -0.784735, 0.372642, + ]; + check_sklearn_parity(4, &train_ref, &probe_ref, &oob_ref, 0.06, 0.12); + } + + #[test] + fn sklearn_parity_two_features() { + // sklearn max_features = 2. Largest std of one run: train 0.0193, probe 0.0137, + // oob 0.0343. + let train_ref = [ + -0.539971, 0.752384, 0.639444, 0.304022, -0.900858, 1.152436, 0.153326, 0.948274, + -0.020744, 0.473991, 0.498276, -0.873499, 0.504429, -1.058965, -0.884246, + -0.399683, -0.943511, 0.897507, -0.102778, 0.483708, 0.157492, 0.602244, 0.124176, + -0.134138, 0.133388, 0.902843, -0.300226, 0.324471, 0.626043, 0.099969, 0.506475, + 0.816064, 0.690784, -0.939260, 0.429617, 0.800137, -0.544917, -0.442264, -0.920062, + 0.086742, + ]; + let probe_ref = [ + 0.819729, 0.481021, -0.705066, -0.572872, 0.188570, -0.723649, -0.302799, 0.719386, + -0.522029, -0.210053, + ]; + let oob_ref = [ + -0.584733, 0.327078, 0.387249, 0.195302, -0.898731, 0.523206, 0.280689, 0.512263, + 0.251790, 0.431326, 0.776490, -0.653868, 0.466336, -0.482361, -0.578605, -0.352370, + -0.705269, 0.764271, -0.065520, 0.412983, -0.046812, 0.460163, 0.044958, 0.472574, + 0.448485, 0.631409, -0.138894, 0.607838, 0.623829, 0.106691, 0.481788, 0.481038, + 0.563559, -0.701948, 0.468302, 0.449402, -0.539726, -0.623343, -0.657852, 0.195843, + ]; + check_sklearn_parity(2, &train_ref, &probe_ref, &oob_ref, 0.08, 0.14); + } + } } From 06a48a71cc9246a855f5e4b46038ef4454cf76f8 Mon Sep 17 00:00:00 2001 From: slievens <5492873+slievens@users.noreply.github.com> Date: Fri, 2 Oct 2026 21:36:57 +0200 Subject: [PATCH 07/17] Tidying up of base_tree_regressor --- src/tree/base_tree_regressor.rs | 97 ++++++++++++++------------------- 1 file changed, 41 insertions(+), 56 deletions(-) diff --git a/src/tree/base_tree_regressor.rs b/src/tree/base_tree_regressor.rs index 1cb7c46f..9f8d2a30 100644 --- a/src/tree/base_tree_regressor.rs +++ b/src/tree/base_tree_regressor.rs @@ -130,10 +130,8 @@ struct NodeVisitor<'a, TX: Number + PartialOrd, TY: Number, X: Array2, Y: Ar x: &'a X, y: &'a Y, node: usize, - // holds the elements for this node, sorted for each feature [num_features, num_samples] - //sorted_node_elements: Vec>, - start_idx: usize, - end_idx: usize, + start_idx: usize, // start index for the elements in `sorted_by_feature` of SplitWorkspace + end_idx: usize, // end index (exclusive) in the same vector(s) true_child_output: f64, false_child_output: f64, level: u16, @@ -175,7 +173,7 @@ fn mass_of(i: usize, samples: &[usize], sample_weights: Option<&[f64]>) -> f64 { } } -// Struct representing an element that belongs logically to a Node, as stored in NodeVisitor +// Struct representing an element that belongs logically to a Node, as stored in SplitWorkspace #[derive(Copy, Clone, Default)] struct NodeElement { // the row index in the dataset @@ -233,22 +231,22 @@ fn stable_partition( w } -struct ScratchPad { - is_true: Vec, - node_elements: Vec, - shared_node_elements: Vec>, +struct SplitWorkspace { + in_true_branch: Vec, // indexed by row index in X + partition_buffer: Vec, + sorted_by_feature: Vec>, } -impl ScratchPad { +impl SplitWorkspace { fn new( - is_true: Vec, - node_elements: Vec, - shared_node_elements: Vec>, + in_true_branch: Vec, + partition_buffer: Vec, + sorted_by_feature: Vec>, ) -> Self { Self { - is_true, - node_elements, - shared_node_elements, + in_true_branch, + partition_buffer, + sorted_by_feature, } } } @@ -319,7 +317,7 @@ impl, Y: Array1> let root = Node::new(sum / mass); nodes.push(root); - let shared_node_elements: Vec> = order + let sorted_by_feature: Vec> = order .iter() .map(|col_order| { col_order @@ -333,12 +331,12 @@ impl, Y: Array1> .collect() }) .collect(); - let end_idx = shared_node_elements[0].len(); + let end_idx = sorted_by_feature[0].len(); - let mut scratch_pad = ScratchPad::new( + let mut workspace = SplitWorkspace::new( vec![false; x.shape().0], vec![NodeElement::default(); end_idx], - shared_node_elements, + sorted_by_feature, ); let mut base_tree = BaseTreeRegressor { @@ -355,14 +353,14 @@ impl, Y: Array1> let mut visitor_queue: VecDeque> = VecDeque::new(); - if base_tree.find_best_cutoff(&mut visitor, mtry, mass, &mut rng, &scratch_pad) { + if base_tree.find_best_cutoff(&mut visitor, mtry, mass, &mut rng, &workspace) { visitor_queue.push_back(visitor); } let max_depth = base_tree.parameters().max_depth.unwrap_or(u16::MAX); while let Some(node) = visitor_queue.pop_front() { if node.level < max_depth { - base_tree.split(node, mtry, &mut visitor_queue, &mut rng, &mut scratch_pad); + base_tree.split(node, mtry, &mut visitor_queue, &mut rng, &mut workspace); } } @@ -407,12 +405,11 @@ impl, Y: Array1> mtry: usize, mass: f64, rng: &mut impl rand::Rng, - scratch_pad: &ScratchPad, + workspace: &SplitWorkspace, ) -> bool { let (_, n_attr) = visitor.x.shape(); - //let n: usize = visitor.num_samples(); - let n: usize = scratch_pad.shared_node_elements[0][visitor.start_idx..visitor.end_idx] + let n: usize = workspace.sorted_by_feature[0][visitor.start_idx..visitor.end_idx] .iter() .map(|elem| elem.count as usize) .sum(); @@ -446,19 +443,11 @@ impl, Y: Array1> parent_gain, *variable, rng, - scratch_pad, + workspace, ); } Splitter::Best => { - self.find_best_split( - visitor, - n, - mass, - sum, - parent_gain, - *variable, - scratch_pad, - ); + self.find_best_split(visitor, n, mass, sum, parent_gain, *variable, workspace); } } } @@ -475,14 +464,14 @@ impl, Y: Array1> parent_gain: f64, j: usize, rng: &mut impl rand::Rng, - scratch_pad: &ScratchPad, + workspace: &SplitWorkspace, ) { if visitor.start_idx == visitor.end_idx { return; } - let first_elem = scratch_pad.shared_node_elements[j][visitor.start_idx]; + let first_elem = workspace.sorted_by_feature[j][visitor.start_idx]; let min_val = visitor.x.get((first_elem.row(), j)); - let last_elem = scratch_pad.shared_node_elements[j][visitor.end_idx - 1]; + let last_elem = workspace.sorted_by_feature[j][visitor.end_idx - 1]; let max_val = visitor.x.get((last_elem.row(), j)); if min_val >= max_val { @@ -494,7 +483,7 @@ impl, Y: Array1> let mut true_sum = 0f64; let mut true_mass = 0f64; let mut true_count = 0; - for elem in &scratch_pad.shared_node_elements[j][visitor.start_idx..visitor.end_idx] { + for elem in &workspace.sorted_by_feature[j][visitor.start_idx..visitor.end_idx] { if visitor.x.get((elem.row(), j)).to_f64().unwrap() <= split_value { true_sum += elem.mass * visitor.y.get(elem.row()).to_f64().unwrap(); true_count += elem.count; @@ -543,14 +532,14 @@ impl, Y: Array1> sum: f64, parent_gain: f64, j: usize, - scratch_pad: &ScratchPad, + workspace: &SplitWorkspace, ) { let mut true_sum = 0f64; let mut true_count = 0; let mut true_mass = 0f64; let mut prevx = Option::None; - for elem in &scratch_pad.shared_node_elements[j][visitor.start_idx..visitor.end_idx] { + for elem in &workspace.sorted_by_feature[j][visitor.start_idx..visitor.end_idx] { let x_ij = *visitor.x.get((elem.row(), j)); if prevx.is_none() || x_ij == prevx.unwrap() { @@ -614,24 +603,20 @@ impl, Y: Array1> mtry: usize, visitor_queue: &mut VecDeque>, rng: &mut impl rand::Rng, - scratch_pad: &mut ScratchPad, // buffer used to track the splitting + workspace: &mut SplitWorkspace, ) -> bool { let this_node = &self.nodes()[visitor.node]; - // sorted_node_elements needs to be turned into - // Vec< (Vec, Vec) > and then into - // (Vec, Vec) - let mut tc = 0usize; let mut true_mass = 0f64; let mut fc = 0usize; let mut false_mass = 0f64; // for each row_index, does it belong in the true branch or not? - let is_true = &mut scratch_pad.is_true; + let in_true_branch = &mut workspace.in_true_branch; let mut n_true = 0usize; - for e in &scratch_pad.shared_node_elements[0][visitor.start_idx..visitor.end_idx] { + for e in &workspace.sorted_by_feature[0][visitor.start_idx..visitor.end_idx] { let t = is_true_sample(e, visitor.x, this_node); - is_true[e.row()] = t; + in_true_branch[e.row()] = t; n_true += t as usize; // Fill in tc, etc while we are at it if t { @@ -666,21 +651,21 @@ impl, Y: Array1> self.depth = u16::max(self.depth, visitor.level + 1); - // If the child nodes can not be split any further, there is no point is partitioning the ranges + // If the child nodes can not be split any further, there is no point in partitioning the ranges let max_depth = self.parameters().max_depth.unwrap_or(u16::MAX); let child_level = visitor.level + 1; let min_split = self.parameters().min_samples_split; let true_can_split = child_level < max_depth && tc >= min_split; let false_can_split = child_level < max_depth && fc >= min_split; if !true_can_split && !false_can_split { - return true; // both children are leaves: no partition, no search + return true; // both children are leaves: no partition } for j in 0..visitor.x.shape().1 { stable_partition( - &mut scratch_pad.shared_node_elements[j][visitor.start_idx..visitor.end_idx], - &mut scratch_pad.node_elements, - is_true, + &mut workspace.sorted_by_feature[j][visitor.start_idx..visitor.end_idx], + &mut workspace.partition_buffer, + in_true_branch, ); } @@ -693,7 +678,7 @@ impl, Y: Array1> visitor.level + 1, ); - if self.find_best_cutoff(&mut true_visitor, mtry, true_mass, rng, scratch_pad) { + if self.find_best_cutoff(&mut true_visitor, mtry, true_mass, rng, workspace) { visitor_queue.push_back(true_visitor); } @@ -706,7 +691,7 @@ impl, Y: Array1> visitor.level + 1, ); - if self.find_best_cutoff(&mut false_visitor, mtry, false_mass, rng, scratch_pad) { + if self.find_best_cutoff(&mut false_visitor, mtry, false_mass, rng, workspace) { visitor_queue.push_back(false_visitor); } From 28ea8b4669dd0153cfe6e318ab97ff1979f3e75b Mon Sep 17 00:00:00 2001 From: slievens <5492873+slievens@users.noreply.github.com> Date: Mon, 5 Oct 2026 16:08:51 +0200 Subject: [PATCH 08/17] Delete test to reduce conflicts --- src/ensemble/random_forest_regressor.rs | 192 ------------------------ 1 file changed, 192 deletions(-) diff --git a/src/ensemble/random_forest_regressor.rs b/src/ensemble/random_forest_regressor.rs index 9ebcc53c..95eeb483 100644 --- a/src/ensemble/random_forest_regressor.rs +++ b/src/ensemble/random_forest_regressor.rs @@ -797,196 +797,4 @@ mod tests { let msg = "'fit' should be called before calling 'predict'"; assert_eq!(yhat.err(), Some(Failed::predict(msg))); } - - mod sklearn_parity { - use super::*; - // sklearn parity tests. - // - // smartcore and numpy use different RNGs, thus the bootstrap samples and the feature - // subsets are different. With many trees, both forests converge to the same bagged - // predictor. Thus we compare the predictions within a tolerance. - // - // Reference: sklearn 1.9.1, numpy 2.4.6. Each reference value is the mean of 10 sklearn - // runs (random_state = 0..10). The tolerance is approximately 4 x the largest standard - // deviation of one sklearn run, for each row. - // - // ```python - // rng = np.random.default_rng(0) - // x = np.round(rng.uniform(-1, 1, (40, 4)), 4) - // y = np.round(x[:, 0] * x[:, 1] + np.sin(3 * x[:, 2]) + 0.1 * rng.normal(size=40), 4) - // x_probe = np.round(rng.uniform(-1, 1, (10, 4)), 4) - // for max_features in [1.0, 2]: - // for seed in range(10): - // rf = RandomForestRegressor(n_estimators=2000, max_features=max_features, - // max_depth=None, min_samples_leaf=1, min_samples_split=2, bootstrap=True, - // oob_score=True, random_state=seed).fit(x, y) - // # collect rf.predict(x), rf.predict(x_probe), rf.oob_prediction_ - // ``` - // - // TODO: no weighted case. smartcore uses the sample weights for the bootstrap and also - // multiplies them into the tree mass. sklearn uses them only for the bootstrap. - - fn sklearn_parity_train_data() -> (DenseMatrix, Vec) { - let x = DenseMatrix::from_2d_array(&[ - &[0.2739, -0.4604, -0.9181, -0.9669], - &[0.6265, 0.8255, 0.2133, 0.459], - &[0.0872, 0.8701, 0.6317, -0.9945], - &[0.7148, -0.9328, 0.4593, -0.6487], - &[0.7264, 0.0829, -0.4006, -0.1546], - &[-0.9434, -0.7514, 0.3412, 0.2944], - &[0.2308, -0.2326, 0.9944, 0.9617], - &[0.3711, 0.3009, 0.3769, -0.2222], - &[-0.7298, 0.443, 0.0507, -0.3795], - &[-0.0283, 0.779, 0.8681, -0.2844], - &[0.1431, -0.3563, 0.1886, -0.3242], - &[-0.2168, 0.7805, -0.5457, 0.2464], - &[-0.832, 0.6653, 0.5742, -0.5213], - &[0.753, -0.8829, -0.3278, -0.6994], - &[-0.0993, 0.5926, -0.5387, -0.896], - &[-0.1909, -0.603, -0.8185, 0.1607], - &[-0.4026, 0.344, -0.601, 0.8842], - &[-0.2698, -0.789, 0.2582, 0.8543], - &[-0.1192, 0.9092, -0.0002, -0.1495], - &[0.2404, 0.9902, 0.8979, -0.0799], - &[0.5155, -0.0052, 0.0586, 0.5716], - &[-0.1707, 0.469, 0.4223, 0.8641], - &[-0.7701, 0.458, 0.8548, 0.9359], - &[-0.9706, 0.7273, 0.9624, 0.9144], - &[-0.7025, 0.9453, 0.7799, 0.6447], - &[-0.04, -0.5353, 0.6038, 0.8471], - &[-0.4677, 0.0779, -0.1145, 0.862], - &[-0.919, 0.464, 0.2287, -0.9433], - &[0.4384, -0.968, 0.5159, 0.0255], - &[0.8582, -0.8678, 0.6826, -0.8666], - &[-0.3114, -0.1394, 0.9321, 0.1245], - &[-0.4823, -0.5166, 0.7762, -0.5483], - &[-0.7509, -0.4233, 0.1722, 0.1082], - &[0.6194, 0.121, -0.4232, -0.1742], - &[0.6362, 0.253, 0.9182, -0.2612], - &[0.1052, 0.1878, 0.6966, -0.7091], - &[-0.187, 0.8199, -0.9139, 0.6454], - &[-0.1692, 0.6596, -0.9801, -0.2699], - &[-0.8427, 0.3052, -0.4523, 0.4053], - &[0.8876, -0.7464, 0.7296, -0.8811], - ]) - .unwrap(); - let y = vec![ - -0.5144, 0.9957, 0.7839, 0.3660, -0.9022, 1.5099, 0.0804, 1.1980, -0.1768, 0.4984, - 0.3364, -1.0023, 0.5267, -1.3905, -1.0531, -0.4267, -1.0746, 0.9736, -0.1242, - 0.5237, 0.2751, 0.6806, 0.1690, -0.4747, -0.0497, 1.0539, -0.3933, 0.1634, 0.6273, - 0.0960, 0.5208, 1.0106, 0.7643, -1.0745, 0.4076, 0.9968, -0.5477, -0.3399, -1.0701, - 0.0243, - ]; - (x, y) - } - - fn sklearn_parity_probe_data() -> DenseMatrix { - DenseMatrix::from_2d_array(&[ - &[-0.3606, -0.625, 0.3451, -0.6098], - &[0.1554, 0.2045, 0.9248, -0.8555], - &[-0.0001, 0.4882, -0.6455, -0.2239], - &[-0.8742, 0.4518, -0.8245, -0.2098], - &[0.747, -0.0554, 0.8252, 0.5318], - &[0.8306, -0.7452, -0.8529, -0.8593], - &[0.7377, 0.2681, -0.0069, -0.6729], - &[0.3475, -0.364, 0.4218, -0.0793], - &[0.0149, 0.5793, -0.8145, 0.1575], - &[-0.6055, 0.6163, -0.0223, 0.9774], - ]) - .unwrap() - } - - fn assert_close_to_sklearn(actual: &[f64], expected: &[f64], tol: f64, label: &str) { - assert_eq!(actual.len(), expected.len(), "{label}: length"); - for (i, (a, e)) in actual.iter().zip(expected.iter()).enumerate() { - assert!( - (a - e).abs() <= tol, - "{label}, row {i}: smartcore {a}, sklearn {e}, tol {tol}" - ); - } - } - - /// Fits the forest with 2000 trees and compares the predictions with sklearn. - fn check_sklearn_parity( - m: usize, - train_ref: &[f64], - probe_ref: &[f64], - oob_ref: &[f64], - tol: f64, - oob_tol: f64, - ) { - let (x, y) = sklearn_parity_train_data(); - let x_probe = sklearn_parity_probe_data(); - - let parameters = RandomForestRegressorParameters::default() - .with_n_trees(2000) - .with_m(m) - .with_min_samples_leaf(1) - .with_min_samples_split(2) - .with_keep_samples(true) - .with_seed(42); - let forest = RandomForestRegressor::fit(&x, &y, parameters).unwrap(); - - let y_hat: Vec = forest.predict(&x).unwrap(); - assert_close_to_sklearn(&y_hat, train_ref, tol, "train"); - - let y_hat_probe: Vec = forest.predict(&x_probe).unwrap(); - assert_close_to_sklearn(&y_hat_probe, probe_ref, tol, "probe"); - - let y_hat_oob: Vec = forest.predict_oob(&x).unwrap(); - assert_close_to_sklearn(&y_hat_oob, oob_ref, oob_tol, "oob"); - } - - #[test] - fn sklearn_parity_all_features() { - // sklearn max_features = 1.0. Largest std of one run: train 0.0132, probe 0.0097, - // oob 0.0290. - let train_ref = [ - -0.563215, 0.790441, 0.671600, 0.447413, -0.962337, 1.182714, 0.146502, 0.956499, - -0.019351, 0.475851, 0.536221, -0.921016, 0.566884, -1.182586, -0.953829, - -0.492116, -0.970510, 0.906254, -0.122531, 0.492432, 0.111251, 0.701316, 0.166136, - -0.090198, 0.168000, 0.917484, -0.318774, 0.452445, 0.701507, 0.213636, 0.519599, - 0.777545, 0.710208, -1.026581, 0.428588, 0.814387, -0.550148, -0.449774, -0.966291, - 0.151257, - ]; - let probe_ref = [ - 0.840299, 0.478916, -0.941862, -0.540659, 0.284103, -0.685080, -0.272447, 0.791806, - -0.562334, -0.195136, - ]; - let oob_ref = [ - -0.649343, 0.431595, 0.475531, 0.590899, -1.064646, 0.606166, 0.261989, 0.535072, - 0.255804, 0.436133, 0.880022, -0.782491, 0.635705, -0.820807, -0.774075, -0.607414, - -0.781085, 0.788450, -0.119615, 0.437044, -0.173755, 0.739059, 0.160968, 0.594925, - 0.542656, 0.672530, -0.189747, 0.961226, 0.830511, 0.416988, 0.517550, 0.376026, - 0.616485, -0.942362, 0.465415, 0.488966, -0.554237, -0.643990, -0.784735, 0.372642, - ]; - check_sklearn_parity(4, &train_ref, &probe_ref, &oob_ref, 0.06, 0.12); - } - - #[test] - fn sklearn_parity_two_features() { - // sklearn max_features = 2. Largest std of one run: train 0.0193, probe 0.0137, - // oob 0.0343. - let train_ref = [ - -0.539971, 0.752384, 0.639444, 0.304022, -0.900858, 1.152436, 0.153326, 0.948274, - -0.020744, 0.473991, 0.498276, -0.873499, 0.504429, -1.058965, -0.884246, - -0.399683, -0.943511, 0.897507, -0.102778, 0.483708, 0.157492, 0.602244, 0.124176, - -0.134138, 0.133388, 0.902843, -0.300226, 0.324471, 0.626043, 0.099969, 0.506475, - 0.816064, 0.690784, -0.939260, 0.429617, 0.800137, -0.544917, -0.442264, -0.920062, - 0.086742, - ]; - let probe_ref = [ - 0.819729, 0.481021, -0.705066, -0.572872, 0.188570, -0.723649, -0.302799, 0.719386, - -0.522029, -0.210053, - ]; - let oob_ref = [ - -0.584733, 0.327078, 0.387249, 0.195302, -0.898731, 0.523206, 0.280689, 0.512263, - 0.251790, 0.431326, 0.776490, -0.653868, 0.466336, -0.482361, -0.578605, -0.352370, - -0.705269, 0.764271, -0.065520, 0.412983, -0.046812, 0.460163, 0.044958, 0.472574, - 0.448485, 0.631409, -0.138894, 0.607838, 0.623829, 0.106691, 0.481788, 0.481038, - 0.563559, -0.701948, 0.468302, 0.449402, -0.539726, -0.623343, -0.657852, 0.195843, - ]; - check_sklearn_parity(2, &train_ref, &probe_ref, &oob_ref, 0.08, 0.14); - } - } } From 219c89d0a07c7792a45880e985d3ac1c1e43e000 Mon Sep 17 00:00:00 2001 From: slievens <5492873+slievens@users.noreply.github.com> Date: Fri, 9 Oct 2026 08:53:08 +0200 Subject: [PATCH 09/17] Do not allocate variables on every iteration --- src/tree/base_tree_regressor.rs | 16 +++++++++------- 1 file changed, 9 insertions(+), 7 deletions(-) diff --git a/src/tree/base_tree_regressor.rs b/src/tree/base_tree_regressor.rs index 9f8d2a30..5fd214bd 100644 --- a/src/tree/base_tree_regressor.rs +++ b/src/tree/base_tree_regressor.rs @@ -235,6 +235,7 @@ struct SplitWorkspace { in_true_branch: Vec, // indexed by row index in X partition_buffer: Vec, sorted_by_feature: Vec>, + variables: Vec, // the variables. Can this be made smaller? } impl SplitWorkspace { @@ -242,11 +243,13 @@ impl SplitWorkspace { in_true_branch: Vec, partition_buffer: Vec, sorted_by_feature: Vec>, + n: usize, ) -> Self { Self { in_true_branch, partition_buffer, sorted_by_feature, + variables: (0..n).collect::>(), } } } @@ -337,6 +340,7 @@ impl, Y: Array1> vec![false; x.shape().0], vec![NodeElement::default(); end_idx], sorted_by_feature, + x.shape().1, ); let mut base_tree = BaseTreeRegressor { @@ -353,7 +357,7 @@ impl, Y: Array1> let mut visitor_queue: VecDeque> = VecDeque::new(); - if base_tree.find_best_cutoff(&mut visitor, mtry, mass, &mut rng, &workspace) { + if base_tree.find_best_cutoff(&mut visitor, mtry, mass, &mut rng, &mut workspace) { visitor_queue.push_back(visitor); } @@ -405,7 +409,7 @@ impl, Y: Array1> mtry: usize, mass: f64, rng: &mut impl rand::Rng, - workspace: &SplitWorkspace, + workspace: &mut SplitWorkspace, ) -> bool { let (_, n_attr) = visitor.x.shape(); @@ -420,11 +424,9 @@ impl, Y: Array1> let sum = self.nodes()[visitor.node].output * mass; - // TODO later: get rid of this allocation in every iteration - let mut variables = (0..n_attr).collect::>(); - + // Note: in sklearn, the attributes are always considered in a random order if mtry < n_attr { - variables.shuffle(rng); + workspace.variables.shuffle(rng); } let parent_gain = @@ -432,7 +434,7 @@ impl, Y: Array1> let splitter = self.parameters().splitter.clone(); - for variable in variables.iter().take(mtry) { + for variable in workspace.variables.iter().take(mtry) { match splitter { Splitter::Random => { self.find_random_split( From db3a0540215dd8b2e00fd960111dfa14bfc23bab Mon Sep 17 00:00:00 2001 From: slievens <5492873+slievens@users.noreply.github.com> Date: Fri, 9 Oct 2026 09:41:06 +0200 Subject: [PATCH 10/17] Make NodeElement as small as possible --- src/tree/base_tree_regressor.rs | 211 ++++++++++++++++++++++++-------- 1 file changed, 160 insertions(+), 51 deletions(-) diff --git a/src/tree/base_tree_regressor.rs b/src/tree/base_tree_regressor.rs index 5fd214bd..c379bb81 100644 --- a/src/tree/base_tree_regressor.rs +++ b/src/tree/base_tree_regressor.rs @@ -173,9 +173,16 @@ fn mass_of(i: usize, samples: &[usize], sample_weights: Option<&[f64]>) -> f64 { } } +trait NodeElement: Copy + Default { + fn new(row: usize, count: usize, mass: f64) -> Self; + fn row(&self) -> usize; + fn count(&self) -> usize; + fn mass(&self) -> f64; +} + // Struct representing an element that belongs logically to a Node, as stored in SplitWorkspace #[derive(Copy, Clone, Default)] -struct NodeElement { +struct WeightedElement { // the row index in the dataset pub row_idx: u32, // the number of times this row is present, should always be > 0 @@ -185,19 +192,98 @@ struct NodeElement { pub mass: f64, } -impl NodeElement { +impl NodeElement for WeightedElement { + fn new(row_idx: usize, count: usize, mass: f64) -> Self { + Self { + row_idx: row_idx as u32, + count: count as u32, + mass, + } + } + // return the row_idx as usize + #[inline(always)] + fn row(&self) -> usize { + self.row_idx as usize + } + + #[inline(always)] + fn count(&self) -> usize { + self.count as usize + } + + #[inline(always)] + fn mass(&self) -> f64 { + self.mass + } +} + +#[derive(Copy, Clone, Default)] +struct CountedElement { + // the row index in the dataset + pub row_idx: u32, + // the number of times this row is present, should always be > 0 + pub count: u32, +} + +impl NodeElement for CountedElement { + fn new(row_idx: usize, count: usize, _mass: f64) -> Self { + Self { + row_idx: row_idx as u32, + count: count as u32, + } + } + // return the row_idx as usize + #[inline(always)] + fn row(&self) -> usize { + self.row_idx as usize + } + + #[inline(always)] + fn count(&self) -> usize { + self.count as usize + } + + #[inline(always)] + fn mass(&self) -> f64 { + self.count as f64 + } +} + +#[derive(Copy, Clone, Default)] +struct UnitElement { + // the row index in the dataset + pub row_idx: u32, +} + +impl NodeElement for UnitElement { + fn new(row_idx: usize, _count: usize, _mass: f64) -> Self { + Self { + row_idx: row_idx as u32, + } + } // return the row_idx as usize #[inline(always)] fn row(&self) -> usize { self.row_idx as usize } + + #[inline(always)] + fn count(&self) -> usize { + 1 + } + + #[inline(always)] + fn mass(&self) -> f64 { + 1.0 + } } // Checks whether the example indicated by node_element is a "true child" for this node -fn is_true_sample(node_element: &NodeElement, x: &X, node: &Node) -> bool +fn is_true_sample(node_element: &E, x: &X, node: &Node) -> bool where TX: Number + PartialOrd, X: Array2, + E: NodeElement, { x.get((node_element.row(), node.split_feature)) .to_f64() @@ -209,18 +295,17 @@ where // scratch: temp buffer // is_true: is_true[idx] checks whether element with row idx equal to idx belongs to the true branch // returns: index of first element of false branch -fn stable_partition( - slice: &mut [NodeElement], - scratch: &mut [NodeElement], - is_true: &[bool], -) -> usize { +fn stable_partition(slice: &mut [E], scratch: &mut [E], is_true: &[bool]) -> usize +where + E: NodeElement, +{ // Note: this is intentionally written without an if/else branch in the main loop let n = slice.len(); let scratch = &mut scratch[..n]; let (mut w, mut f) = (0usize, 0usize); for i in 0..n { let e = slice[i]; - let t = is_true[e.row_idx as usize]; + let t = is_true[e.row()]; slice[w] = e; // w <= i, so this never clobbers an unread element scratch[f] = e; // advance only one of the pointers @@ -231,18 +316,18 @@ fn stable_partition( w } -struct SplitWorkspace { +struct SplitWorkspace { in_true_branch: Vec, // indexed by row index in X - partition_buffer: Vec, - sorted_by_feature: Vec>, + partition_buffer: Vec, + sorted_by_feature: Vec>, variables: Vec, // the variables. Can this be made smaller? } -impl SplitWorkspace { +impl SplitWorkspace { fn new( in_true_branch: Vec, - partition_buffer: Vec, - sorted_by_feature: Vec>, + partition_buffer: Vec, + sorted_by_feature: Vec>, n: usize, ) -> Self { Self { @@ -302,6 +387,34 @@ impl, Y: Array1> mtry: usize, order: &[Vec], parameters: BaseTreeRegressorParameters, + ) -> Result, Failed> { + match sample_weights { + Some(_) => Self::grow::( + x, + y, + sample_weights, + samples, + mtry, + order, + parameters, + ), + None if samples.iter().all(|&s| s == 1) => { + Self::grow::(x, y, sample_weights, samples, mtry, order, parameters) + } + None => { + Self::grow::(x, y, sample_weights, samples, mtry, order, parameters) + } + } + } + + fn grow( + x: &X, + y: &Y, + sample_weights: Option<&[f64]>, + samples: Vec, + mtry: usize, + order: &[Vec], + parameters: BaseTreeRegressorParameters, ) -> Result, Failed> { let n_rows = y.shape(); @@ -320,17 +433,13 @@ impl, Y: Array1> let root = Node::new(sum / mass); nodes.push(root); - let sorted_by_feature: Vec> = order + let sorted_by_feature: Vec> = order .iter() .map(|col_order| { col_order .iter() .filter(|&&i| samples[i] > 0) - .map(|&i| NodeElement { - row_idx: i as u32, - count: samples[i] as u32, - mass: mass_of(i, &samples, sample_weights), - }) + .map(|&i| E::new(i, samples[i], mass_of(i, &samples, sample_weights))) .collect() }) .collect(); @@ -338,7 +447,7 @@ impl, Y: Array1> let mut workspace = SplitWorkspace::new( vec![false; x.shape().0], - vec![NodeElement::default(); end_idx], + vec![E::default(); end_idx], sorted_by_feature, x.shape().1, ); @@ -403,19 +512,19 @@ impl, Y: Array1> } } - fn find_best_cutoff( + fn find_best_cutoff( &mut self, visitor: &mut NodeVisitor<'_, TX, TY, X, Y>, mtry: usize, mass: f64, rng: &mut impl rand::Rng, - workspace: &mut SplitWorkspace, + workspace: &mut SplitWorkspace, ) -> bool { let (_, n_attr) = visitor.x.shape(); let n: usize = workspace.sorted_by_feature[0][visitor.start_idx..visitor.end_idx] .iter() - .map(|elem| elem.count as usize) + .map(|elem| elem.count()) .sum(); if n < self.parameters().min_samples_split { @@ -457,7 +566,7 @@ impl, Y: Array1> self.nodes()[visitor.node].split_score.is_some() } - fn find_random_split( + fn find_random_split( &mut self, visitor: &mut NodeVisitor<'_, TX, TY, X, Y>, n: usize, @@ -466,7 +575,7 @@ impl, Y: Array1> parent_gain: f64, j: usize, rng: &mut impl rand::Rng, - workspace: &SplitWorkspace, + workspace: &SplitWorkspace, ) { if visitor.start_idx == visitor.end_idx { return; @@ -487,15 +596,15 @@ impl, Y: Array1> let mut true_count = 0; for elem in &workspace.sorted_by_feature[j][visitor.start_idx..visitor.end_idx] { if visitor.x.get((elem.row(), j)).to_f64().unwrap() <= split_value { - true_sum += elem.mass * visitor.y.get(elem.row()).to_f64().unwrap(); - true_count += elem.count; - true_mass += elem.mass; + true_sum += elem.mass() * visitor.y.get(elem.row()).to_f64().unwrap(); + true_count += elem.count(); + true_mass += elem.mass(); } } - let false_count = n - (true_count as usize); + let false_count = n - true_count; - if (true_count as usize) < self.parameters().min_samples_leaf + if true_count < self.parameters().min_samples_leaf || false_count < self.parameters().min_samples_leaf { return; @@ -526,7 +635,7 @@ impl, Y: Array1> } } - fn find_best_split( + fn find_best_split( &mut self, visitor: &mut NodeVisitor<'_, TX, TY, X, Y>, n: usize, @@ -534,7 +643,7 @@ impl, Y: Array1> sum: f64, parent_gain: f64, j: usize, - workspace: &SplitWorkspace, + workspace: &SplitWorkspace, ) { let mut true_sum = 0f64; let mut true_count = 0; @@ -546,21 +655,21 @@ impl, Y: Array1> if prevx.is_none() || x_ij == prevx.unwrap() { prevx = Some(x_ij); - true_count += elem.count; - true_mass += elem.mass; - true_sum += elem.mass * visitor.y.get(elem.row()).to_f64().unwrap(); + true_count += elem.count(); + true_mass += elem.mass(); + true_sum += elem.mass() * visitor.y.get(elem.row()).to_f64().unwrap(); continue; } - let false_count = n - (true_count as usize); + let false_count = n - true_count; - if (true_count as usize) < self.parameters().min_samples_leaf + if true_count < self.parameters().min_samples_leaf || false_count < self.parameters().min_samples_leaf { prevx = Some(x_ij); - true_count += elem.count; - true_mass += elem.mass; - true_sum += elem.mass * visitor.y.get(elem.row()).to_f64().unwrap(); + true_count += elem.count(); + true_mass += elem.mass(); + true_sum += elem.mass() * visitor.y.get(elem.row()).to_f64().unwrap(); continue; } @@ -593,19 +702,19 @@ impl, Y: Array1> } prevx = Some(x_ij); - true_sum += elem.mass * visitor.y.get(elem.row()).to_f64().unwrap(); - true_count += elem.count; - true_mass += elem.mass; + true_sum += elem.mass() * visitor.y.get(elem.row()).to_f64().unwrap(); + true_count += elem.count(); + true_mass += elem.mass(); } } - fn split<'a>( + fn split<'a, E: NodeElement>( &mut self, visitor: NodeVisitor<'a, TX, TY, X, Y>, mtry: usize, visitor_queue: &mut VecDeque>, rng: &mut impl rand::Rng, - workspace: &mut SplitWorkspace, + workspace: &mut SplitWorkspace, ) -> bool { let this_node = &self.nodes()[visitor.node]; @@ -622,11 +731,11 @@ impl, Y: Array1> n_true += t as usize; // Fill in tc, etc while we are at it if t { - tc += e.count as usize; - true_mass += e.mass; + tc += e.count(); + true_mass += e.mass(); } else { - fc += e.count as usize; - false_mass += e.mass; + fc += e.count(); + false_mass += e.mass(); } } From 927d0e81793b32fc1455d9b6b9de6fd55f396f9a Mon Sep 17 00:00:00 2001 From: slievens <5492873+slievens@users.noreply.github.com> Date: Fri, 9 Oct 2026 10:22:35 +0200 Subject: [PATCH 11/17] Small cleanup of base_tree_regressor --- src/tree/base_tree_regressor.rs | 71 ++++++++++++++++++--------------- 1 file changed, 39 insertions(+), 32 deletions(-) diff --git a/src/tree/base_tree_regressor.rs b/src/tree/base_tree_regressor.rs index c379bb81..80a30034 100644 --- a/src/tree/base_tree_regressor.rs +++ b/src/tree/base_tree_regressor.rs @@ -1,5 +1,4 @@ use std::collections::VecDeque; -use std::default::Default; use std::fmt::Debug; use std::marker::PhantomData; @@ -173,34 +172,34 @@ fn mass_of(i: usize, samples: &[usize], sample_weights: Option<&[f64]>) -> f64 { } } +// Trait representing an element that belongs logically to a Node. Implementations of it are stored in SplitWorkspace trait NodeElement: Copy + Default { fn new(row: usize, count: usize, mass: f64) -> Self; + // The row index in the dataset fn row(&self) -> usize; + // The number of times this row was sampled, should always be > 0 fn count(&self) -> usize; + // The total mass of this element: count * sample_weight of row or count if no sample weights are used fn mass(&self) -> f64; } -// Struct representing an element that belongs logically to a Node, as stored in SplitWorkspace #[derive(Copy, Clone, Default)] struct WeightedElement { - // the row index in the dataset - pub row_idx: u32, - // the number of times this row is present, should always be > 0 - pub count: u32, - // total mass of this element, equals count * mass of individual element. - // equals count when no sample weights were used - pub mass: f64, + row_idx: u32, + count: u32, + mass: f64, } impl NodeElement for WeightedElement { fn new(row_idx: usize, count: usize, mass: f64) -> Self { + debug_assert!(count > 0); Self { row_idx: row_idx as u32, count: count as u32, mass, } } - // return the row_idx as usize + #[inline(always)] fn row(&self) -> usize { self.row_idx as usize @@ -219,20 +218,19 @@ impl NodeElement for WeightedElement { #[derive(Copy, Clone, Default)] struct CountedElement { - // the row index in the dataset - pub row_idx: u32, - // the number of times this row is present, should always be > 0 - pub count: u32, + row_idx: u32, + count: u32, } impl NodeElement for CountedElement { fn new(row_idx: usize, count: usize, _mass: f64) -> Self { + debug_assert!(count > 0); Self { row_idx: row_idx as u32, count: count as u32, } } - // return the row_idx as usize + #[inline(always)] fn row(&self) -> usize { self.row_idx as usize @@ -251,17 +249,17 @@ impl NodeElement for CountedElement { #[derive(Copy, Clone, Default)] struct UnitElement { - // the row index in the dataset - pub row_idx: u32, + row_idx: u32, } impl NodeElement for UnitElement { - fn new(row_idx: usize, _count: usize, _mass: f64) -> Self { + fn new(row_idx: usize, count: usize, _mass: f64) -> Self { + debug_assert_eq!(count, 1); Self { row_idx: row_idx as u32, } } - // return the row_idx as usize + #[inline(always)] fn row(&self) -> usize { self.row_idx as usize @@ -320,7 +318,7 @@ struct SplitWorkspace { in_true_branch: Vec, // indexed by row index in X partition_buffer: Vec, sorted_by_feature: Vec>, - variables: Vec, // the variables. Can this be made smaller? + variables: Vec, // feature indices which will be shuffled when mtry < n_features } impl SplitWorkspace { @@ -398,12 +396,12 @@ impl, Y: Array1> order, parameters, ), - None if samples.iter().all(|&s| s == 1) => { - Self::grow::(x, y, sample_weights, samples, mtry, order, parameters) - } - None => { - Self::grow::(x, y, sample_weights, samples, mtry, order, parameters) + None if samples.iter().all(|&s| s <= 1) => { + // Note: if any sample has count zero, it will be filtered out, so we should + // use the smallest possible NodeElement + Self::grow::(x, y, None, samples, mtry, order, parameters) } + None => Self::grow::(x, y, None, samples, mtry, order, parameters), } } @@ -718,9 +716,9 @@ impl, Y: Array1> ) -> bool { let this_node = &self.nodes()[visitor.node]; - let mut tc = 0usize; + let mut true_count = 0usize; let mut true_mass = 0f64; - let mut fc = 0usize; + let mut false_count = 0usize; let mut false_mass = 0f64; // for each row_index, does it belong in the true branch or not? let in_true_branch = &mut workspace.in_true_branch; @@ -731,16 +729,18 @@ impl, Y: Array1> n_true += t as usize; // Fill in tc, etc while we are at it if t { - tc += e.count(); + true_count += e.count(); true_mass += e.mass(); } else { - fc += e.count(); + false_count += e.count(); false_mass += e.mass(); } } // Stop early if it is clear that there will be too few examples in the leaf - if tc < self.parameters().min_samples_leaf || fc < self.parameters().min_samples_leaf { + if true_count < self.parameters().min_samples_leaf + || false_count < self.parameters().min_samples_leaf + { self.nodes[visitor.node].split_feature = 0; self.nodes[visitor.node].split_value = Option::None; self.nodes[visitor.node].split_score = Option::None; @@ -766,8 +766,8 @@ impl, Y: Array1> let max_depth = self.parameters().max_depth.unwrap_or(u16::MAX); let child_level = visitor.level + 1; let min_split = self.parameters().min_samples_split; - let true_can_split = child_level < max_depth && tc >= min_split; - let false_can_split = child_level < max_depth && fc >= min_split; + let true_can_split = child_level < max_depth && true_count >= min_split; + let false_can_split = child_level < max_depth && false_count >= min_split; if !true_can_split && !false_can_split { return true; // both children are leaves: no partition } @@ -1184,4 +1184,11 @@ mod tests { let y_hat_repeated = tree_repeated.predict(&x).expect("Predict should work"); assert!(mean_absolute_error(&y_hat, &y_hat_repeated) < 1e-9); } + + #[test] + fn test_node_element_sizes() { + assert_eq!(size_of::(), 4); + assert_eq!(size_of::(), 8); + assert_eq!(size_of::(), 16); + } } From 9fe19cd4a8167a0bb459dfeb5a10b22722a3d557 Mon Sep 17 00:00:00 2001 From: slievens <5492873+slievens@users.noreply.github.com> Date: Fri, 9 Oct 2026 12:17:19 +0200 Subject: [PATCH 12/17] Improve setup cost: less alloc, less branching --- src/tree/base_tree_regressor.rs | 24 +++++++++++++++--------- 1 file changed, 15 insertions(+), 9 deletions(-) diff --git a/src/tree/base_tree_regressor.rs b/src/tree/base_tree_regressor.rs index 80a30034..67680e46 100644 --- a/src/tree/base_tree_regressor.rs +++ b/src/tree/base_tree_regressor.rs @@ -192,7 +192,6 @@ struct WeightedElement { impl NodeElement for WeightedElement { fn new(row_idx: usize, count: usize, mass: f64) -> Self { - debug_assert!(count > 0); Self { row_idx: row_idx as u32, count: count as u32, @@ -224,7 +223,6 @@ struct CountedElement { impl NodeElement for CountedElement { fn new(row_idx: usize, count: usize, _mass: f64) -> Self { - debug_assert!(count > 0); Self { row_idx: row_idx as u32, count: count as u32, @@ -253,8 +251,7 @@ struct UnitElement { } impl NodeElement for UnitElement { - fn new(row_idx: usize, count: usize, _mass: f64) -> Self { - debug_assert_eq!(count, 1); + fn new(row_idx: usize, _count: usize, _mass: f64) -> Self { Self { row_idx: row_idx as u32, } @@ -431,16 +428,25 @@ impl, Y: Array1> let root = Node::new(sum / mass); nodes.push(root); + let counts: Vec = samples.iter().map(|&s| s as u32).collect(); + // number of distinct rows of x in this tree + let n_kept = counts.iter().filter(|&&c| c > 0).count(); + let sorted_by_feature: Vec> = order .iter() .map(|col_order| { - col_order - .iter() - .filter(|&&i| samples[i] > 0) - .map(|&i| E::new(i, samples[i], mass_of(i, &samples, sample_weights))) - .collect() + let mut out = vec![E::default(); n_kept + 1]; // preallocate, one additional place + let mut w = 0usize; // index to write to + for &i in col_order { + let c = counts[i]; + out[w] = E::new(i, c as usize, mass_of(i, &samples, sample_weights)); + w += (c > 0) as usize; // update w in a branchless way + } + out.truncate(n_kept); + out }) .collect(); + let end_idx = sorted_by_feature[0].len(); let mut workspace = SplitWorkspace::new( From 6284d67ee8b0734fbb5c1ef038d832a581225f88 Mon Sep 17 00:00:00 2001 From: slievens <5492873+slievens@users.noreply.github.com> Date: Fri, 9 Oct 2026 12:53:18 +0200 Subject: [PATCH 13/17] Inline(never) + stack instead of queue --- src/tree/base_tree_regressor.rs | 3 ++- 1 file changed, 2 insertions(+), 1 deletion(-) diff --git a/src/tree/base_tree_regressor.rs b/src/tree/base_tree_regressor.rs index 67680e46..6617a45c 100644 --- a/src/tree/base_tree_regressor.rs +++ b/src/tree/base_tree_regressor.rs @@ -290,6 +290,7 @@ where // scratch: temp buffer // is_true: is_true[idx] checks whether element with row idx equal to idx belongs to the true branch // returns: index of first element of false branch +#[inline(never)] fn stable_partition(slice: &mut [E], scratch: &mut [E], is_true: &[bool]) -> usize where E: NodeElement, @@ -475,7 +476,7 @@ impl, Y: Array1> } let max_depth = base_tree.parameters().max_depth.unwrap_or(u16::MAX); - while let Some(node) = visitor_queue.pop_front() { + while let Some(node) = visitor_queue.pop_back() { if node.level < max_depth { base_tree.split(node, mtry, &mut visitor_queue, &mut rng, &mut workspace); } From f3fe872af34913e19e2ab1e455cec92726e284a3 Mon Sep 17 00:00:00 2001 From: slievens <5492873+slievens@users.noreply.github.com> Date: Fri, 9 Oct 2026 16:27:03 +0200 Subject: [PATCH 14/17] Rewrite CountedElement as one u64 --- src/tree/base_tree_regressor.rs | 20 ++++++-------------- 1 file changed, 6 insertions(+), 14 deletions(-) diff --git a/src/tree/base_tree_regressor.rs b/src/tree/base_tree_regressor.rs index 6617a45c..e596b319 100644 --- a/src/tree/base_tree_regressor.rs +++ b/src/tree/base_tree_regressor.rs @@ -215,33 +215,25 @@ impl NodeElement for WeightedElement { } } +// implement CountedElement as a single u64, so it can be retrieved in a single read #[derive(Copy, Clone, Default)] -struct CountedElement { - row_idx: u32, - count: u32, -} +struct CountedElement(u64); // low 32 bits: row, high 32 bits: count impl NodeElement for CountedElement { fn new(row_idx: usize, count: usize, _mass: f64) -> Self { - Self { - row_idx: row_idx as u32, - count: count as u32, - } + Self(row_idx as u64 | ((count as u64) << 32)) } - #[inline(always)] fn row(&self) -> usize { - self.row_idx as usize + self.0 as u32 as usize } - #[inline(always)] fn count(&self) -> usize { - self.count as usize + (self.0 >> 32) as usize } - #[inline(always)] fn mass(&self) -> f64 { - self.count as f64 + (self.0 >> 32) as f64 } } From 8e8826e79b12dc9d7c15200543af0475b6042f5d Mon Sep 17 00:00:00 2001 From: slievens <5492873+slievens@users.noreply.github.com> Date: Fri, 9 Oct 2026 17:48:28 +0200 Subject: [PATCH 15/17] Final cleanup. - Use plain vec instead of vecdeque for stack. - Loop over the split feature to determine true/false elements. This saves some lookups. --- src/ensemble/base_forest_regressor.rs | 2 +- src/tree/base_tree_regressor.rs | 95 +++++++++++++++------------ 2 files changed, 55 insertions(+), 42 deletions(-) diff --git a/src/ensemble/base_forest_regressor.rs b/src/ensemble/base_forest_regressor.rs index 8970c64c..b48f8dad 100644 --- a/src/ensemble/base_forest_regressor.rs +++ b/src/ensemble/base_forest_regressor.rs @@ -148,7 +148,7 @@ impl, Y: Array1 min_samples_leaf: parameters.min_samples_leaf, min_samples_split: parameters.min_samples_split, seed: Some(parameters.seed.wrapping_add(tree_idx as u64)), // give each tree its own fixed seed - splitter: parameters.splitter.clone(), + splitter: parameters.splitter, }; // Only use sample weights on base tree if not already applied during bootstrapping let sample_weights_for_base_tree = if parameters.bootstrap { diff --git a/src/tree/base_tree_regressor.rs b/src/tree/base_tree_regressor.rs index e596b319..e2eff66f 100644 --- a/src/tree/base_tree_regressor.rs +++ b/src/tree/base_tree_regressor.rs @@ -1,4 +1,3 @@ -use std::collections::VecDeque; use std::fmt::Debug; use std::marker::PhantomData; @@ -14,7 +13,7 @@ use crate::numbers::basenum::Number; use crate::rand_custom::get_rng_impl; #[cfg_attr(feature = "serde", derive(Serialize, Deserialize))] -#[derive(Debug, Clone, Default)] +#[derive(Debug, Clone, Copy, Default)] pub enum Splitter { Random, #[default] @@ -114,10 +113,10 @@ impl, Y: Array1> PartialE for BaseTreeRegressor { fn eq(&self, other: &Self) -> bool { - if self.depth != other.depth || self.nodes().len() != other.nodes().len() { + if self.depth != other.depth || self.nodes.len() != other.nodes.len() { false } else { - self.nodes() + self.nodes .iter() .zip(other.nodes().iter()) .all(|(a, b)| a == b) @@ -316,13 +315,13 @@ impl SplitWorkspace { in_true_branch: Vec, partition_buffer: Vec, sorted_by_feature: Vec>, - n: usize, + n_features: usize, ) -> Self { Self { in_true_branch, partition_buffer, sorted_by_feature, - variables: (0..n).collect::>(), + variables: (0..n_features).collect::>(), } } } @@ -461,16 +460,16 @@ impl, Y: Array1> let mut visitor = NodeVisitor::::new(0, 0, end_idx, x, y, 1); - let mut visitor_queue: VecDeque> = VecDeque::new(); + let mut visitor_stack: Vec> = Vec::new(); if base_tree.find_best_cutoff(&mut visitor, mtry, mass, &mut rng, &mut workspace) { - visitor_queue.push_back(visitor); + visitor_stack.push(visitor); } let max_depth = base_tree.parameters().max_depth.unwrap_or(u16::MAX); - while let Some(node) = visitor_queue.pop_back() { + while let Some(node) = visitor_stack.pop() { if node.level < max_depth { - base_tree.split(node, mtry, &mut visitor_queue, &mut rng, &mut workspace); + base_tree.split(node, mtry, &mut visitor_stack, &mut rng, &mut workspace); } } @@ -494,7 +493,7 @@ impl, Y: Array1> pub(crate) fn predict_for_row(&self, x: &X, row: usize) -> TY { let mut node_id = 0; loop { - let node = &self.nodes()[node_id]; + let node = &self.nodes[node_id]; let Some(true_child) = node.true_child else { return TY::from_f64(node.output).unwrap(); }; @@ -528,17 +527,16 @@ impl, Y: Array1> return false; } - let sum = self.nodes()[visitor.node].output * mass; + let sum = self.nodes[visitor.node].output * mass; // Note: in sklearn, the attributes are always considered in a random order if mtry < n_attr { workspace.variables.shuffle(rng); } - let parent_gain = - mass * self.nodes()[visitor.node].output * self.nodes()[visitor.node].output; + let parent_gain = mass * self.nodes[visitor.node].output * self.nodes[visitor.node].output; - let splitter = self.parameters().splitter.clone(); + let splitter = self.parameters().splitter; for variable in workspace.variables.iter().take(mtry) { match splitter { @@ -560,7 +558,7 @@ impl, Y: Array1> } } - self.nodes()[visitor.node].split_score.is_some() + self.nodes[visitor.node].split_score.is_some() } fn find_random_split( @@ -685,8 +683,8 @@ impl, Y: Array1> let gain = (true_mass * true_mean * true_mean + false_mass * false_mean * false_mean) - parent_gain; - if self.nodes()[visitor.node].split_score.is_none() - || gain > self.nodes()[visitor.node].split_score.unwrap() + if self.nodes[visitor.node].split_score.is_none() + || gain > self.nodes[visitor.node].split_score.unwrap() { self.nodes[visitor.node].split_feature = j; self.nodes[visitor.node].split_value = Option::Some( @@ -705,15 +703,21 @@ impl, Y: Array1> } } + /// Apply the split that was found for `visitor.node`: add its two child nodes and + /// partition the node's range in each `sorted_by_feature` column into a true part + /// and a false part. Then find the best cutoff for each child. If a child can be + /// split, push it on `visitor_stack`. + /// If a branch has fewer than `min_samples_leaf` samples, clear the split and + /// keep the node as a leaf. fn split<'a, E: NodeElement>( &mut self, visitor: NodeVisitor<'a, TX, TY, X, Y>, mtry: usize, - visitor_queue: &mut VecDeque>, + visitor_stack: &mut Vec>, rng: &mut impl rand::Rng, workspace: &mut SplitWorkspace, ) -> bool { - let this_node = &self.nodes()[visitor.node]; + let this_node = &self.nodes[visitor.node]; let mut true_count = 0usize; let mut true_mass = 0f64; @@ -721,19 +725,24 @@ impl, Y: Array1> let mut false_mass = 0f64; // for each row_index, does it belong in the true branch or not? let in_true_branch = &mut workspace.in_true_branch; + + // the column for split_feature is in order, so all the true samples should come first, all the false samples second + let split_feature = this_node.split_feature; + let col = &workspace.sorted_by_feature[split_feature][visitor.start_idx..visitor.end_idx]; let mut n_true = 0usize; - for e in &workspace.sorted_by_feature[0][visitor.start_idx..visitor.end_idx] { - let t = is_true_sample(e, visitor.x, this_node); - in_true_branch[e.row()] = t; - n_true += t as usize; - // Fill in tc, etc while we are at it - if t { - true_count += e.count(); - true_mass += e.mass(); - } else { - false_count += e.count(); - false_mass += e.mass(); + for e in col { + if !is_true_sample(e, visitor.x, this_node) { + break; } + in_true_branch[e.row()] = true; + true_count += e.count(); + true_mass += e.mass(); + n_true += 1; + } + for e in &col[n_true..] { + in_true_branch[e.row()] = false; + false_count += e.count(); + false_mass += e.mass(); } // Stop early if it is clear that there will be too few examples in the leaf @@ -750,20 +759,20 @@ impl, Y: Array1> // Add the child nodes to the tree let split_idx = visitor.start_idx + n_true; - let true_child_idx = self.nodes().len(); + let true_child_idx = self.nodes.len(); self.nodes.push(Node::new(visitor.true_child_output)); - let false_child_idx = self.nodes().len(); + let false_child_idx = self.nodes.len(); self.nodes.push(Node::new(visitor.false_child_output)); self.nodes[visitor.node].true_child = Some(true_child_idx); self.nodes[visitor.node].false_child = Some(false_child_idx); - self.depth = u16::max(self.depth, visitor.level + 1); + let child_level = visitor.level + 1; + self.depth = u16::max(self.depth, child_level); // If the child nodes can not be split any further, there is no point in partitioning the ranges let max_depth = self.parameters().max_depth.unwrap_or(u16::MAX); - let child_level = visitor.level + 1; let min_split = self.parameters().min_samples_split; let true_can_split = child_level < max_depth && true_count >= min_split; let false_can_split = child_level < max_depth && false_count >= min_split; @@ -785,11 +794,13 @@ impl, Y: Array1> split_idx, visitor.x, visitor.y, - visitor.level + 1, + child_level, ); - if self.find_best_cutoff(&mut true_visitor, mtry, true_mass, rng, workspace) { - visitor_queue.push_back(true_visitor); + if true_can_split + && self.find_best_cutoff(&mut true_visitor, mtry, true_mass, rng, workspace) + { + visitor_stack.push(true_visitor); } let mut false_visitor = NodeVisitor::::new( @@ -798,11 +809,13 @@ impl, Y: Array1> visitor.end_idx, visitor.x, visitor.y, - visitor.level + 1, + child_level, ); - if self.find_best_cutoff(&mut false_visitor, mtry, false_mass, rng, workspace) { - visitor_queue.push_back(false_visitor); + if false_can_split + && self.find_best_cutoff(&mut false_visitor, mtry, false_mass, rng, workspace) + { + visitor_stack.push(false_visitor); } true From d071328c00b58e1d0f52253bee0c1239f337c847 Mon Sep 17 00:00:00 2001 From: slievens <5492873+slievens@users.noreply.github.com> Date: Fri, 9 Oct 2026 20:48:30 +0200 Subject: [PATCH 16/17] Return error for too many rows and test empty-X fit - Add check_row_count to BaseTreeRegressor::fit_weak_learner. It returns a Failed error when x has more rows than u32::MAX, because node elements store row indices and counts as u32. - Stop the random-split scan at the first element above the split value. The elements are sorted by the feature, so no later element can be in the true branch. - Add tests: check_row_count limits, and fit with an empty x (0x0, 0x3, 3x0) returns an error for the decision tree, random forest and extra trees regressors. - Use std::mem::size_of in test_node_element_sizes and make comments clearer. --- src/ensemble/extra_trees_regressor.rs | 25 +++++++++++++++++ src/ensemble/random_forest_regressor.rs | 25 +++++++++++++++++ src/tree/base_tree_regressor.rs | 36 +++++++++++++++++++++---- src/tree/decision_tree_regressor.rs | 22 +++++++++++++++ 4 files changed, 103 insertions(+), 5 deletions(-) diff --git a/src/ensemble/extra_trees_regressor.rs b/src/ensemble/extra_trees_regressor.rs index c0dd1a99..a5a06f77 100644 --- a/src/ensemble/extra_trees_regressor.rs +++ b/src/ensemble/extra_trees_regressor.rs @@ -277,6 +277,7 @@ impl, Y: Array1 #[cfg(test)] mod tests { use super::*; + use crate::error::FailedError; use crate::linalg::basic::matrix::DenseMatrix; use crate::metrics::mean_squared_error; @@ -470,6 +471,30 @@ mod tests { assert_eq!(yhat.err(), Some(Failed::predict(msg))); } + #[cfg_attr( + all(target_arch = "wasm32", not(target_os = "wasi")), + wasm_bindgen_test::wasm_bindgen_test + )] + #[test] + fn fit_with_empty_x_should_not_panic() { + let parameters = ExtraTreesRegressorParameters::default() + .with_n_trees(5) + .with_seed(42); + // (rows, columns, y): no rows, no columns, or both + let cases: Vec<(usize, usize, Vec)> = + vec![(0, 0, vec![]), (0, 3, vec![]), (3, 0, vec![1.0, 2.0, 3.0])]; + for (nrows, ncols, y) in cases { + let x: DenseMatrix = DenseMatrix::new(nrows, ncols, vec![], false) + .expect("Construction of empty x should work"); + let result = ExtraTreesRegressor::fit(&x, &y, parameters.clone()); + let expected = Failed::because( + FailedError::ParametersError, + "Training data must contain at least one sample and one feature.", + ); + assert_eq!(result.err(), Some(expected), "shape: ({nrows}, {ncols})"); + } + } + mod sklearn_parity { use super::*; // sklearn parity tests. diff --git a/src/ensemble/random_forest_regressor.rs b/src/ensemble/random_forest_regressor.rs index 9dabdf2d..364806e4 100644 --- a/src/ensemble/random_forest_regressor.rs +++ b/src/ensemble/random_forest_regressor.rs @@ -470,6 +470,7 @@ impl, Y: Array1 #[cfg(test)] mod tests { use super::*; + use crate::error::FailedError; use crate::linalg::basic::matrix::DenseMatrix; use crate::metrics::mean_absolute_error; @@ -801,6 +802,30 @@ mod tests { assert_eq!(yhat.err(), Some(Failed::predict(msg))); } + #[cfg_attr( + all(target_arch = "wasm32", not(target_os = "wasi")), + wasm_bindgen_test::wasm_bindgen_test + )] + #[test] + fn fit_with_empty_x_should_not_panic() { + let parameters = RandomForestRegressorParameters::default() + .with_n_trees(5) + .with_seed(42); + // (rows, columns, y): no rows, no columns, or both + let cases: Vec<(usize, usize, Vec)> = + vec![(0, 0, vec![]), (0, 3, vec![]), (3, 0, vec![1.0, 2.0, 3.0])]; + for (nrows, ncols, y) in cases { + let x: DenseMatrix = DenseMatrix::new(nrows, ncols, vec![], false) + .expect("Construction of empty x should work"); + let result = RandomForestRegressor::fit(&x, &y, parameters.clone()); + let expected = Failed::because( + FailedError::ParametersError, + "Training data must contain at least one sample and one feature.", + ); + assert_eq!(result.err(), Some(expected), "shape: ({nrows}, {ncols})"); + } + } + mod sklearn_parity { use super::*; // sklearn parity tests. diff --git a/src/tree/base_tree_regressor.rs b/src/tree/base_tree_regressor.rs index e2eff66f..8ff9959a 100644 --- a/src/tree/base_tree_regressor.rs +++ b/src/tree/base_tree_regressor.rs @@ -171,6 +171,16 @@ fn mass_of(i: usize, samples: &[usize], sample_weights: Option<&[f64]>) -> f64 { } } +/// Node elements store row indices and counts as `u32`. +fn check_row_count(n_rows: usize) -> Result<(), Failed> { + if u32::try_from(n_rows).is_err() { + return Err(Failed::fit( + "Number of rows in x must not be larger than u32::MAX", + )); + } + Ok(()) +} + // Trait representing an element that belongs logically to a Node. Implementations of it are stored in SplitWorkspace trait NodeElement: Copy + Default { fn new(row: usize, count: usize, mass: f64) -> Self; @@ -281,12 +291,12 @@ where // scratch: temp buffer // is_true: is_true[idx] checks whether element with row idx equal to idx belongs to the true branch // returns: index of first element of false branch -#[inline(never)] +#[inline(never)] // inline(never) is faster as seen during profiling fn stable_partition(slice: &mut [E], scratch: &mut [E], is_true: &[bool]) -> usize where E: NodeElement, { - // Note: this is intentionally written without an if/else branch in the main loop + // Note: this is intentionally written without an explicit if/else branch in the main loop let n = slice.len(); let scratch = &mut scratch[..n]; let (mut w, mut f) = (0usize, 0usize); @@ -375,6 +385,7 @@ impl, Y: Array1> order: &[Vec], parameters: BaseTreeRegressorParameters, ) -> Result, Failed> { + check_row_count(y.shape())?; match sample_weights { Some(_) => Self::grow::( x, @@ -594,6 +605,8 @@ impl, Y: Array1> true_sum += elem.mass() * visitor.y.get(elem.row()).to_f64().unwrap(); true_count += elem.count(); true_mass += elem.mass(); + } else { + break; } } @@ -829,6 +842,19 @@ mod tests { use crate::linalg::basic::matrix::DenseMatrix; use crate::metrics::mean_absolute_error; + #[test] + #[cfg(target_pointer_width = "64")] + fn check_row_count_rejects_more_rows_than_u32_max() { + assert!(check_row_count(0).is_ok()); + assert!(check_row_count(u32::MAX as usize).is_ok()); + assert_eq!( + check_row_count(u32::MAX as usize + 1).err(), + Some(Failed::fit( + "Number of rows in x must not be larger than u32::MAX" + )) + ); + } + #[test] fn test_fit_on_empty_data_returns_error() { // 2 rows x 2 features — values are arbitrary; only the empty-row case is under test @@ -1199,8 +1225,8 @@ mod tests { #[test] fn test_node_element_sizes() { - assert_eq!(size_of::(), 4); - assert_eq!(size_of::(), 8); - assert_eq!(size_of::(), 16); + assert_eq!(std::mem::size_of::(), 4); + assert_eq!(std::mem::size_of::(), 8); + assert_eq!(std::mem::size_of::(), 16); } } diff --git a/src/tree/decision_tree_regressor.rs b/src/tree/decision_tree_regressor.rs index 839cfdac..12ab72b6 100644 --- a/src/tree/decision_tree_regressor.rs +++ b/src/tree/decision_tree_regressor.rs @@ -370,6 +370,7 @@ impl, Y: Array1> #[cfg(test)] mod tests { use super::*; + use crate::error::FailedError; use crate::linalg::basic::matrix::DenseMatrix; #[test] @@ -629,6 +630,27 @@ mod tests { assert_eq!(yhat.err(), Some(Failed::predict(msg))); } + #[cfg_attr( + all(target_arch = "wasm32", not(target_os = "wasi")), + wasm_bindgen_test::wasm_bindgen_test + )] + #[test] + fn fit_with_empty_x_should_not_panic() { + // (rows, columns, y): no rows, no columns, or both + let cases: Vec<(usize, usize, Vec)> = + vec![(0, 0, vec![]), (0, 3, vec![]), (3, 0, vec![1.0, 2.0, 3.0])]; + for (nrows, ncols, y) in cases { + let x: DenseMatrix = DenseMatrix::new(nrows, ncols, vec![], false) + .expect("Construction of empty x should work"); + let result = DecisionTreeRegressor::fit(&x, &y, Default::default()); + let expected = Failed::because( + FailedError::ParametersError, + "Training data must contain at least one sample and one feature.", + ); + assert_eq!(result.err(), Some(expected), "shape: ({nrows}, {ncols})"); + } + } + #[cfg_attr( all(target_arch = "wasm32", not(target_os = "wasi")), wasm_bindgen_test::wasm_bindgen_test From 201efd8e26402a55935f113678a8408d156fc686 Mon Sep 17 00:00:00 2001 From: slievens <5492873+slievens@users.noreply.github.com> Date: Sat, 10 Oct 2026 14:34:16 +0200 Subject: [PATCH 17/17] Make tree split thresholds consistent with f64 routing Split search compared feature values in TX, but partitioning and predict compare them as f64 against an f64 threshold. For large integers and adjacent or very large f64 values, the applied split was then different from the split that was scored. - find_best_split: detect ties on f64 values. Compute the midpoint as prev + (x - prev) / 2. Use prev as the threshold if the midpoint rounds up to x or overflows. - find_random_split: convert the bounds to f64 before the range check, so random_range does not get an empty range and panic. - CountedElement/UnitElement: add debug_asserts that sample weights are not used, because their mass() ignores weights. - Add tests for i64 values above 2^53, adjacent f64 values and f64 midpoint overflow. --- src/tree/base_tree_regressor.rs | 121 ++++++++++++++++++++++++++++++-- 1 file changed, 114 insertions(+), 7 deletions(-) diff --git a/src/tree/base_tree_regressor.rs b/src/tree/base_tree_regressor.rs index 8ff9959a..1bdf4fad 100644 --- a/src/tree/base_tree_regressor.rs +++ b/src/tree/base_tree_regressor.rs @@ -229,7 +229,9 @@ impl NodeElement for WeightedElement { struct CountedElement(u64); // low 32 bits: row, high 32 bits: count impl NodeElement for CountedElement { - fn new(row_idx: usize, count: usize, _mass: f64) -> Self { + fn new(row_idx: usize, count: usize, mass: f64) -> Self { + // mass() is derived from count, so sample weights must not be used + debug_assert_eq!(mass, count as f64); Self(row_idx as u64 | ((count as u64) << 32)) } #[inline(always)] @@ -252,7 +254,11 @@ struct UnitElement { } impl NodeElement for UnitElement { - fn new(row_idx: usize, _count: usize, _mass: f64) -> Self { + fn new(row_idx: usize, count: usize, mass: f64) -> Self { + // count() and mass() are always 1, so sample weights and counts > 1 must not be used. + // A count of 0 is allowed: such elements are overwritten during setup. + debug_assert!(count <= 1); + debug_assert_eq!(mass, count as f64); Self { row_idx: row_idx as u32, } @@ -591,11 +597,16 @@ impl, Y: Array1> let last_elem = workspace.sorted_by_feature[j][visitor.end_idx - 1]; let max_val = visitor.x.get((last_elem.row(), j)); + // Convert min_val and max_val to f64 first, so that values that are different in type TX + // but equal in f64 do not end up panicking the call to the rng.random_range function. + let min_val = min_val.to_f64().unwrap(); + let max_val = max_val.to_f64().unwrap(); + if min_val >= max_val { return; } - let split_value = rng.random_range(min_val.to_f64().unwrap()..max_val.to_f64().unwrap()); + let split_value = rng.random_range(min_val..max_val); let mut true_sum = 0f64; let mut true_mass = 0f64; @@ -659,7 +670,9 @@ impl, Y: Array1> let mut prevx = Option::None; for elem in &workspace.sorted_by_feature[j][visitor.start_idx..visitor.end_idx] { - let x_ij = *visitor.x.get((elem.row(), j)); + // Comparison for equality is done directly on f64 to avoid later problems + // with values being unequal in TX but equal in f64. + let x_ij = visitor.x.get((elem.row(), j)).to_f64().unwrap(); if prevx.is_none() || x_ij == prevx.unwrap() { prevx = Some(x_ij); @@ -699,10 +712,17 @@ impl, Y: Array1> if self.nodes[visitor.node].split_score.is_none() || gain > self.nodes[visitor.node].split_score.unwrap() { + let prev = prevx.unwrap(); + // compute the midpoint in a way that does not overflow when x_ij and prevx have the same sign + let mid = prev + (x_ij - prev) / 2.0; + + // The midpoint should put prevx in the true branch, and x_ij in the false branch. + // However, mid could round up to x_ij, in which case we would also assign x_ij to the true branch. + // In this case, we set the midpoint equal to prevx, which is always smaller + // than x_ij. + let threshold = if mid < x_ij { mid } else { prev }; self.nodes[visitor.node].split_feature = j; - self.nodes[visitor.node].split_value = Option::Some( - (x_ij.to_f64().unwrap() + prevx.unwrap().to_f64().unwrap()) / 2f64, - ); + self.nodes[visitor.node].split_value = Option::Some(threshold); self.nodes[visitor.node].split_score = Option::Some(gain); visitor.true_child_output = true_mean; @@ -1223,6 +1243,93 @@ mod tests { assert!(mean_absolute_error(&y_hat, &y_hat_repeated) < 1e-9); } + // Above 2^53, f64 cannot hold every integer. Split search compares values in TX, + // but partitioning and predict compare them as f64 against an f64 threshold. + const TWO_POW_53: i64 = 1 << 53; + + fn split_value_parameters(splitter: Splitter) -> BaseTreeRegressorParameters { + BaseTreeRegressorParameters { + max_depth: None, + min_samples_leaf: 1, + min_samples_split: 2, + seed: Some(42), + splitter, + } + } + + #[test] + fn large_integer_midpoint_threshold_separates_values() { + // Both values are exact in f64, but the midpoint 2^53 + 3 is not. + // It rounds to 2^53 + 4 (ties to even), so both rows go to the true branch. + let x = + DenseMatrix::from_2d_vec(&vec![vec![TWO_POW_53 + 2], vec![TWO_POW_53 + 4]]).unwrap(); + let y = vec![0.0f64, 10.0]; + + let tree = + BaseTreeRegressor::fit_inner(&x, &y, None, split_value_parameters(Splitter::Best)) + .expect("Fit should work"); + + let threshold = tree.nodes()[0].split_value.expect("Root should split"); + assert!((TWO_POW_53 + 2) as f64 <= threshold); + assert!(((TWO_POW_53 + 4) as f64) > threshold); + + let y_hat = tree.predict(&x).expect("Predict should work"); + assert_eq!(y_hat, vec![0.0, 10.0]); + } + + #[test] + fn large_integer_values_equal_as_f64_do_not_split() { + // 2^53 and 2^53 + 1 differ as i64 but are the same f64. No f64 threshold can + // separate them, so the tree must not split and must predict the mean. + let x = DenseMatrix::from_2d_vec(&vec![vec![TWO_POW_53], vec![TWO_POW_53 + 1]]).unwrap(); + let y = vec![0.0f64, 10.0]; + + let tree = + BaseTreeRegressor::fit_inner(&x, &y, None, split_value_parameters(Splitter::Best)) + .expect("Fit should work"); + + let y_hat = tree.predict(&x).expect("Predict should work"); + assert_eq!(y_hat, vec![5.0, 5.0]); + } + + #[test] + fn large_integer_values_equal_as_f64_random_splitter() { + // min_val < max_val as i64, but the f64 range is empty. This must not panic. + let x = DenseMatrix::from_2d_vec(&vec![vec![TWO_POW_53], vec![TWO_POW_53 + 1]]).unwrap(); + let y = vec![0.0f64, 10.0]; + + let tree = + BaseTreeRegressor::fit_inner(&x, &y, None, split_value_parameters(Splitter::Random)) + .expect("Fit should work"); + + let y_hat = tree.predict(&x).expect("Predict should work"); + assert_eq!(y_hat, vec![5.0, 5.0]); + } + + #[test] + fn adjacent_f64_midpoint_threshold_separates_values() { + // a and b are adjacent f64 values. Their exact midpoint 1 + 1.5 * EPSILON is not + // representable. It is a tie, and ties round to the even mantissa, which is b. + // So the threshold becomes b, and both rows go to the true branch. + let a = 1.0 + f64::EPSILON; + let b = 1.0 + 2.0 * f64::EPSILON; + assert_eq!((a + b) / 2.0, b); // the rounding that causes the problem + + let x = DenseMatrix::from_2d_vec(&vec![vec![a], vec![b]]).unwrap(); + let y = vec![0.0f64, 10.0]; + + let tree = + BaseTreeRegressor::fit_inner(&x, &y, None, split_value_parameters(Splitter::Best)) + .expect("Fit should work"); + + let threshold = tree.nodes()[0].split_value.expect("Root should split"); + assert!(a <= threshold); + assert!(b > threshold); + + let y_hat = tree.predict(&x).expect("Predict should work"); + assert_eq!(y_hat, vec![0.0, 10.0]); + } + #[test] fn test_node_element_sizes() { assert_eq!(std::mem::size_of::(), 4);