diff --git a/src/ensemble/base_forest_regressor.rs b/src/ensemble/base_forest_regressor.rs index f00681ec..550d600e 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, @@ -476,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, @@ -512,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 @@ -527,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/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..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 @@ -797,4 +800,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()), + ); + } + } }