From 6490ee04ccbf1ea0f4882ca4ec54e5ebe6a43ef8 Mon Sep 17 00:00:00 2001 From: slievens <5492873+slievens@users.noreply.github.com> Date: Sun, 4 Oct 2026 18:25:14 +0200 Subject: [PATCH 1/2] Fix: behavior of sample weights. Match the behavior for sklearn on sample weights. - when bootstrap is true, sample weights are only used during the bootstrap phase - when bootstrap is false, then sample weights are passed down to individual trees. Added tests for this: Compare RandomForestRegressor and ExtraTreesRegressor predictions (train, probe, OOB) with reference values from sklearn 1.9.1. Each reference value is the mean of 10 sklearn runs. smartcore and numpy use different RNGs, so the tests use a tolerance of about 4 x the standard deviation of one sklearn run. For RandomForestRegressor: cover max_features = all and 2, with and without sample weights. For ExtraTreesRegressor: cover max_features = all, with and without sample weights. --- src/ensemble/base_forest_regressor.rs | 8 +- src/ensemble/extra_trees_regressor.rs | 210 ++++++++++++++++++ src/ensemble/random_forest_regressor.rs | 274 ++++++++++++++++++++++++ 3 files changed, 491 insertions(+), 1 deletion(-) diff --git a/src/ensemble/base_forest_regressor.rs b/src/ensemble/base_forest_regressor.rs index f00681ec..6571b5d4 100644 --- a/src/ensemble/base_forest_regressor.rs +++ b/src/ensemble/base_forest_regressor.rs @@ -141,10 +141,16 @@ impl, Y: Array1 seed: Some(parameters.seed.wrapping_add(tree_idx as u64)), // give each tree its own fixed seed splitter: parameters.splitter.clone(), }; + // Only use sample weights on base tree if not already applied during bootstrapping + let sample_weights_for_base_tree = if parameters.bootstrap { + None + } else { + sample_weights + }; let tree = BaseTreeRegressor::fit_weak_learner( x, y, - sample_weights, + sample_weights_for_base_tree, samples.clone(), mtry, params, diff --git a/src/ensemble/extra_trees_regressor.rs b/src/ensemble/extra_trees_regressor.rs index f445e3bf..c0dd1a99 100644 --- a/src/ensemble/extra_trees_regressor.rs +++ b/src/ensemble/extra_trees_regressor.rs @@ -469,4 +469,214 @@ 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. + // However, 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]: + // train, probe = [], [] + // for seed in range(N_SEEDS): + // rf = ExtraTreesRegressor( + // n_estimators=N_TREES, + // max_features=max_features, + // max_depth=None, + // min_samples_leaf=5, # so that the predictions are not equal to the training values + // min_samples_split=2, + // bootstrap=False, + // random_state=seed, + // n_jobs=-1, + // ).fit(x, y) + // train.append(rf.predict(x)) + // probe.append(rf.predict(x_probe)) + // print(f"\nExtra Trees. === max_features={max_features}") + // for name, runs in [("train", train), ("probe", probe)]: + // runs = np.array(runs) + // mean = runs.mean(axis=0) + // max_std = runs.std(axis=0, ddof=1).max() + // max_dev = np.abs(runs - mean).max() + // print(f"{name}: max_std={max_std:.5f} max_dev={max_dev:.5f}") + // print(f"{name}_ref:", rust_vec(mean)) + // ``` + // + + 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 sklearn_sample_weights() -> Vec { + (0..40).into_iter().map(|i| ((i % 4) + 1) as f64).collect() + } + + 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 extra trees regressor with 2000 trees and compares the predictions with sklearn. + fn check_sklearn_parity( + m: usize, + train_ref: &[f64], + probe_ref: &[f64], + tol: f64, + sample_weights: Option<&[f64]>, + ) { + let (x, y) = sklearn_parity_train_data(); + let x_probe = sklearn_parity_probe_data(); + + // Use `min_samples_leaf`= 5 to prevent the tree from predicting all training examples correct + let parameters = ExtraTreesRegressorParameters::default() + .with_n_trees(2000) + .with_m(m) + .with_min_samples_leaf(5) + .with_min_samples_split(2) + .with_keep_samples(true) + .with_seed(42); + let forest = match sample_weights { + None => ExtraTreesRegressor::fit(&x, &y, parameters).unwrap(), + Some(sample_weights) => { + ExtraTreesRegressor::fit_with_weights(&x, &y, sample_weights, 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"); + } + + #[test] + fn sklearn_parity_all_features() { + /* + * train: max_std=0.01081 max_dev=0.02053 + * probe: max_std=0.01127 max_dev=0.02510 + */ + let train_ref = [ + -0.514400, 0.350098, 0.500618, 0.370298, -0.582169, 0.630296, 0.362747, 0.528070, + 0.027928, 0.409103, 0.376290, -0.672295, 0.408925, -0.406530, -0.656713, -0.433254, + -0.656863, 0.572460, -0.005012, 0.398251, 0.134855, 0.457975, 0.232323, 0.203198, + 0.255614, 0.580654, -0.171093, 0.266703, 0.538009, 0.339515, 0.425172, 0.512873, + 0.452168, -0.597154, 0.382057, 0.522851, -0.636264, -0.652577, -0.614220, 0.327194, + ]; + let probe_ref = [ + 0.538464, 0.434785, -0.641271, -0.638714, 0.324906, -0.499588, -0.042825, 0.545153, + -0.637955, -0.073064, + ]; + check_sklearn_parity(4, &train_ref, &probe_ref, 0.045, None); + } + + #[test] + fn sklearn_parity_with_weights() { + /* + * train: max_std=0.01332 max_dev=0.02559 + * probe: max_std=0.01166 max_dev=0.01867 + */ + let train_ref = [ + -0.494622, 0.335181, 0.532771, 0.398018, -0.535612, 0.545873, 0.337625, 0.588692, + 0.030596, 0.431604, 0.376213, -0.686286, 0.400965, -0.400551, -0.658719, -0.481374, + -0.649386, 0.477041, -0.017966, 0.424634, 0.126455, 0.419039, 0.196019, 0.148182, + 0.238049, 0.512727, -0.195074, 0.255755, 0.513608, 0.384085, 0.438159, 0.544918, + 0.374897, -0.552449, 0.433162, 0.583000, -0.636756, -0.643531, -0.627507, 0.372639, + ]; + + let probe_ref = [ + 0.530162, 0.489443, -0.611812, -0.626195, 0.352540, -0.502707, 0.001104, 0.551536, + -0.621779, -0.120032, + ]; + + check_sklearn_parity( + 4, + &train_ref, + &probe_ref, + 0.05, + Some(&sklearn_sample_weights()), + ); + } + } } diff --git a/src/ensemble/random_forest_regressor.rs b/src/ensemble/random_forest_regressor.rs index 95eeb483..6334a1b5 100644 --- a/src/ensemble/random_forest_regressor.rs +++ b/src/ensemble/random_forest_regressor.rs @@ -797,4 +797,278 @@ 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_ + // ``` + // + + 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 sklearn_sample_weights() -> Vec { + (0..40).into_iter().map(|i| ((i % 4) + 1) as f64).collect() + } + + 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, + sample_weights: Option<&[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 = match sample_weights { + None => RandomForestRegressor::fit(&x, &y, parameters).unwrap(), + Some(sample_weights) => { + RandomForestRegressor::fit_with_weights(&x, &y, sample_weights, 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, None); + } + + #[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, None); + } + + #[test] + fn sklearn_parity_with_weights() { + /* train: max_std=0.01455 max_dev=0.03734 + * probe: max_std=0.01280 max_dev=0.02387 + * oob: max_std=0.03646 max_dev=0.07464 + */ + let train_ref = [ + -0.579075, 0.757098, 0.681036, 0.422236, -0.966882, 1.086633, 0.123891, 1.069910, + 0.012417, 0.497026, 0.505539, -0.949550, 0.585693, -1.119453, -0.980359, -0.468459, + -0.905580, 0.871368, -0.134145, 0.509760, 0.091451, 0.714195, 0.177695, -0.283949, + 0.336420, 0.896071, -0.324312, 0.298485, 0.762092, 0.175884, 0.548944, 0.884472, + 0.648075, -1.009675, 0.436657, 0.911007, -0.571345, -0.498412, -0.983307, 0.096344, + ]; + let probe_ref = [ + 0.894188, 0.477767, -0.950106, -0.579763, 0.357390, -0.654198, -0.218833, 0.869894, + -0.585196, -0.204378, + ]; + let oob_ref = [ + -0.610721, 0.467067, 0.441084, 0.654051, -0.999940, 0.545649, 0.227121, 0.542453, + 0.106474, 0.495335, 0.906098, -0.727918, 0.614949, -0.778016, -0.804240, -0.643123, + -0.821945, 0.744144, -0.158162, 0.453301, -0.000066, 0.756773, 0.198170, 0.497569, + 0.524053, 0.703722, -0.164544, 0.868727, 0.828938, 0.275783, 0.615440, 0.383522, + 0.591260, -0.928361, 0.504742, 0.564432, -0.583036, -0.695436, -0.786008, 0.396527, + ]; + + check_sklearn_parity( + 4, + &train_ref, + &probe_ref, + &oob_ref, + 0.06, + 0.20, // set a little high + Some(&sklearn_sample_weights()), + ); + } + + #[test] + fn sklearn_parity_with_weights_two_features() { + /* + * train: max_std=0.01686 max_dev=0.03435 + * probe: max_std=0.01493 max_dev=0.03213 + * oob: max_std=0.05181 max_dev=0.08933 + */ + let train_ref = [ + -0.550424, 0.716045, 0.669342, 0.340627, -0.835593, 1.050114, 0.146434, 1.065546, + 0.013482, 0.465432, 0.480859, -0.924076, 0.454231, -0.991295, -0.909932, -0.425557, + -0.850313, 0.847289, -0.111871, 0.509527, 0.165458, 0.572463, 0.135126, -0.302609, + 0.223818, 0.860349, -0.291573, 0.228765, 0.650378, 0.094705, 0.515905, 0.908350, + 0.658607, -0.858620, 0.452023, 0.898403, -0.559861, -0.491019, -0.931630, 0.063906, + ]; + let probe_ref = [ + 0.852481, 0.503229, -0.635851, -0.585194, 0.276241, -0.719231, -0.182757, 0.771101, + -0.488099, -0.207485, + ]; + let oob_ref = [ + -0.568080, 0.376018, 0.402009, 0.235654, -0.801487, 0.462343, 0.303037, 0.520005, + 0.108155, 0.424214, 0.823031, -0.595168, 0.418338, -0.488094, -0.562944, -0.421132, + -0.739436, 0.690178, -0.082967, 0.451987, 0.110805, 0.435434, 0.055336, 0.401684, + 0.356611, 0.624397, -0.056188, 0.503624, 0.661848, 0.093107, 0.503889, 0.501511, + 0.606863, -0.587745, 0.555991, 0.500660, -0.565841, -0.678327, -0.617068, 0.229146, + ]; + + check_sklearn_parity( + 2, + &train_ref, + &probe_ref, + &oob_ref, + 0.064, + 0.2, + Some(&sklearn_sample_weights()), + ); + } + } } From f3e7271e64a9862a0f915b5d84817cf71617fb6e Mon Sep 17 00:00:00 2001 From: slievens <5492873+slievens@users.noreply.github.com> Date: Mon, 5 Oct 2026 15:49:00 +0200 Subject: [PATCH 2/2] Fix failing tests due to change in sample weights --- src/ensemble/base_forest_regressor.rs | 13 ++++++++----- src/ensemble/random_forest_regressor.rs | 7 +++++-- 2 files changed, 13 insertions(+), 7 deletions(-) diff --git a/src/ensemble/base_forest_regressor.rs b/src/ensemble/base_forest_regressor.rs index 6571b5d4..550d600e 100644 --- a/src/ensemble/base_forest_regressor.rs +++ b/src/ensemble/base_forest_regressor.rs @@ -482,12 +482,12 @@ mod tests { let sample_weights: Vec = (0..20).map(|i| if i < 10 { 1.0 } else { 9.0 }).collect(); let parameters = BaseForestRegressorParameters { - max_depth: Some(1), + max_depth: Some(0), min_samples_leaf: 1, min_samples_split: 2, n_trees: 10, m: None, - keep_samples: true, // keep samples used for each tree, so we can check that they are different + keep_samples: true, seed: 42, bootstrap: false, // No bootstrapping splitter: crate::tree::base_tree_regressor::Splitter::Best, @@ -518,7 +518,7 @@ mod tests { // Use bootstrapping let parameters = BaseForestRegressorParameters { - max_depth: Some(1), + max_depth: Some(0), // Match sibling test on RandomForestRegressor min_samples_leaf: 1, min_samples_split: 2, n_trees: 500, // Use more trees than before to smooth out randomness @@ -533,9 +533,12 @@ mod tests { .expect("Fit should work"); let y_hat = forest.predict(&x).expect("Predict should work"); - // weighted mean is (10.0 * 0 + 90*10) / 100 = 9 + // with bootstrapping the weight should be close to 9, but not extremely close for p in y_hat.iter() { - assert!(p > &9.0f64, "expected value well above 9, got {p}"); + assert!( + (p - 9.0).abs() < 0.1, + "expected value reasonably close to 9, got {p}" + ); } // Without weights, the predicted value should be reasonably close to 5 diff --git a/src/ensemble/random_forest_regressor.rs b/src/ensemble/random_forest_regressor.rs index 6334a1b5..9dabdf2d 100644 --- a/src/ensemble/random_forest_regressor.rs +++ b/src/ensemble/random_forest_regressor.rs @@ -680,9 +680,12 @@ mod tests { .expect("Fit should work"); let y_hat = forest.predict(&x).expect("Predict should work"); - // Due to bootstrapping, weighted value should be well above 9. + // with bootstrapping the weight should be close to 9, but not extremely close for p in y_hat.iter() { - assert!(p > &9.0f64, "expected value greater than 9, got {p}"); + assert!( + (p - 9.0).abs() < 0.1, + "expected value reasonably close to 9, got {p}" + ); } // Without weights, the predicted value should be close to 5