Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
21 changes: 15 additions & 6 deletions src/ensemble/base_forest_regressor.rs
Original file line number Diff line number Diff line change
Expand Up @@ -141,10 +141,16 @@ impl<TX: Number + FloatNumber + PartialOrd, TY: Number, X: Array2<TX>, 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,
Expand Down Expand Up @@ -476,12 +482,12 @@ mod tests {
let sample_weights: Vec<f64> = (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,
Expand Down Expand Up @@ -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
Expand All @@ -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
Expand Down
210 changes: 210 additions & 0 deletions src/ensemble/extra_trees_regressor.rs
Original file line number Diff line number Diff line change
Expand Up @@ -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<f64>, Vec<f64>) {
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<f64> {
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<f64> {
(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<f64> = forest.predict(&x).unwrap();
assert_close_to_sklearn(&y_hat, train_ref, tol, "train");

let y_hat_probe: Vec<f64> = 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()),
);
}
}
}
Loading
Loading