diff --git a/src/ensemble/base_forest_regressor.rs b/src/ensemble/base_forest_regressor.rs index 768f64c4..f00681ec 100644 --- a/src/ensemble/base_forest_regressor.rs +++ b/src/ensemble/base_forest_regressor.rs @@ -120,7 +120,7 @@ impl, Y: Array1 }) .transpose()?; - for _ in 0..parameters.n_trees { + for tree_idx in 0..parameters.n_trees { if parameters.bootstrap { samples = BaseForestRegressor::::sample_with_replacement( n_rows, @@ -138,7 +138,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(parameters.seed.wrapping_add(tree_idx as u64)), // give each tree its own fixed seed splitter: parameters.splitter.clone(), }; let tree = BaseTreeRegressor::fit_weak_learner( @@ -429,6 +429,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 @@ -507,4 +544,67 @@ mod tests { ); } } + + #[test] + fn fit_twice_with_same_seed_gives_identical_forest() { + // With bootstrap, m = 1 and the Random splitter, each random path is used: + // bootstrap sample, feature selection and random thresholds. + 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(); + let sample_weights: Vec = (0..n_rows).map(|i| 1.0 + (i % 4) as f64).collect(); + + for splitter in [Splitter::Best, Splitter::Random] { + for weights in [None, Some(sample_weights.as_slice())] { + let params = BaseForestRegressorParameters { + max_depth: None, + min_samples_leaf: 1, + min_samples_split: 2, + n_trees: 10, + m: Some(1), + keep_samples: true, + seed: 42, + bootstrap: true, + splitter: splitter.clone(), + }; + + let forest_a = BaseForestRegressor::fit(&x, &y, weights, params.clone()) + .expect("Fit should work"); + let forest_b = BaseForestRegressor::fit(&x, &y, weights, params.clone()) + .expect("Fit should work"); + + assert_eq!( + forest_a, forest_b, + "forests differ (splitter: {splitter:?}, weights: {weights:?})" + ); + assert_eq!( + forest_a.samples, forest_b.samples, + "bootstrap samples differ (splitter: {splitter:?}, weights: {weights:?})" + ); + assert_eq!( + forest_a.predict(&x).unwrap(), + forest_b.predict(&x).unwrap(), + "predictions differ (splitter: {splitter:?}, weights: {weights:?})" + ); + + // A different seed must give a different forest, else the check above is trivial + let forest_c = BaseForestRegressor::fit( + &x, + &y, + weights, + BaseForestRegressorParameters { seed: 43, ..params }, + ) + .expect("Fit should work"); + assert_ne!( + forest_a, forest_c, + "seeds 42 and 43 give the same forest (splitter: {splitter:?}, weights: {weights:?})" + ); + } + } + } } diff --git a/src/ensemble/random_forest_classifier.rs b/src/ensemble/random_forest_classifier.rs index ff7a4329..cd26426a 100644 --- a/src/ensemble/random_forest_classifier.rs +++ b/src/ensemble/random_forest_classifier.rs @@ -491,7 +491,7 @@ impl, Y: Array1 = RandomForestClassifier::::sample_with_replacement(&yi, k, &mut rng); if let Some(ref mut all_samples) = maybe_all_samples { @@ -503,7 +503,7 @@ impl, Y: Array1 = DenseMatrix::from_iterator( + (0..n_cols * n_rows).map(|k| ((k / n_cols) * (k % n_cols + 1)) as f64), + n_rows, + n_cols, + 0, + ); + let y: Vec = (0..n_rows) + .map(|i| if i < n_rows / 2 { 0 } else { 1 }) + .collect(); + + let params = RandomForestClassifierParameters { + criterion: SplitCriterion::Gini, + max_depth: Some(2), + min_samples_leaf: 1, + min_samples_split: 2, + n_trees: 10, + m: Some(1), + keep_samples: false, + seed: 42, + }; + let forest = RandomForestClassifier::fit(&x, &y, params).unwrap(); + + // Each tree has one split, thus one feature has a positive importance + let root_features: Vec = forest + .trees + .unwrap() + .iter() + .map(|tree| { + tree.compute_feature_importances(false) + .iter() + .position(|importance| *importance > 0.0) + .unwrap_or(usize::MAX) + }) + .collect(); + + assert!( + root_features.iter().any(|f| *f != root_features[0]), + "all trees split on the same feature: {root_features:?}" + ); + } }