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
104 changes: 102 additions & 2 deletions src/ensemble/base_forest_regressor.rs
Original file line number Diff line number Diff line change
Expand Up @@ -120,7 +120,7 @@ impl<TX: Number + FloatNumber + PartialOrd, TY: Number, X: Array2<TX>, Y: Array1
})
.transpose()?;

for _ in 0..parameters.n_trees {
for tree_idx in 0..parameters.n_trees {
if parameters.bootstrap {
samples = BaseForestRegressor::<TX, TY, X, Y>::sample_with_replacement(
n_rows,
Expand All @@ -138,7 +138,7 @@ impl<TX: Number + FloatNumber + PartialOrd, TY: Number, X: Array2<TX>, 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(
Expand Down Expand Up @@ -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<f64> = DenseMatrix::from_iterator(
(0..4 * n_rows).map(|k| ((k * 7919) % 101) as f64),
n_rows,
4,
0,
);
let y: Vec<f64> = (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
Expand Down Expand Up @@ -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<f64> = DenseMatrix::from_iterator(
(0..4 * n_rows).map(|k| ((k * 7919) % 101) as f64),
n_rows,
4,
0,
);
let y: Vec<f64> = (0..n_rows).map(|i| ((i * 31) % 17) as f64).collect();
let sample_weights: Vec<f64> = (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:?})"
);
}
}
}
}
55 changes: 53 additions & 2 deletions src/ensemble/random_forest_classifier.rs
Original file line number Diff line number Diff line change
Expand Up @@ -491,7 +491,7 @@ impl<TX: FloatNumber + PartialOrd, TY: Number + Ord, X: Array2<TX>, Y: Array1<TY
maybe_all_samples = Some(Vec::with_capacity(n_trees));
}

for _ in 0..parameters.n_trees {
for tree_idx in 0..parameters.n_trees {
let samples: Vec<usize> =
RandomForestClassifier::<TX, TY, X, Y>::sample_with_replacement(&yi, k, &mut rng);
if let Some(ref mut all_samples) = maybe_all_samples {
Expand All @@ -503,7 +503,7 @@ impl<TX: FloatNumber + PartialOrd, TY: Number + Ord, X: Array2<TX>, Y: Array1<TY
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
};
let tree = DecisionTreeClassifier::fit_weak_learner(x, y, samples, mtry, params)?;
trees.push(tree);
Expand Down Expand Up @@ -942,4 +942,55 @@ mod tests {

assert_eq!(forest, deserialized_forest);
}

#[test]
fn each_tree_gets_different_feature_sample() {
// The classifier always uses bootstrapping, thus the trees are different also when
// they all use the same feature sample. For this reason, we look at the features
// directly. If all trees use the same seed, all trees split on the
// same feature.

// Create some data: x[i][j] = i * (j + 1)
let n_rows = 30;
let n_cols = 4;
let x: DenseMatrix<f64> = 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<u32> = (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<usize> = 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:?}"
);
}
}
Loading