feat(tree,ensemble): add sample weights to regressor fit - #467
Conversation
Add `fit_with_weights` to DecisionTreeRegressor, RandomForestRegressor and ExtraTreesRegressor. `fit` keeps its signature and delegates with no weights. - BaseTreeRegressor: use weighted mass (count * weight) for node outputs, split means and gain in the best and random splitters. Return an error when the weight length does not match x. - BaseForestRegressor: take `sample_weights` in `fit`. Bootstrap draws use `WeightedIndex` when weights are given. Pass weights to each tree. - Add tests: weighted root mean, uniform weights equal no weights, integer weights equal repeated rows, weighted full depth, forest balance property, distinct bootstrap samples per tree, seeded determinism, sample_weights properties
Codecov Report❌ Patch coverage is Additional details and impacted files@@ Coverage Diff @@
## main #467 +/- ##
===========================================
+ Coverage 43.97% 63.88% +19.91%
===========================================
Files 85 96 +11
Lines 7281 8413 +1132
===========================================
+ Hits 3202 5375 +2173
+ Misses 4079 3038 -1041 ☔ View full report in Codecov by Harness. 🚀 New features to boost your workflow:
|
Mec-iS
left a comment
There was a problem hiding this comment.
Overall, this is a solid, thoughtful implementation. The approach of adding fit_with_weights alongside fit (which delegates with None) is the right API pattern for backward compatibility. The test suite is impressively comprehensive. Below are the issues I found, ordered by severity.
🔴 Correctness Issues
1. find_best_split: true_mass not reset between candidate splits
In find_best_split, true_mass is updated inside the loop but is never reset when moving to the next split candidate — unlike true_count and true_sum in the original code. Concretely, true_mass is not incremented in the earlier continue branches (for prevx.is_none() || x_ij == prevx.unwrap() and the min_samples_leaf guard). This means true_mass and true_sum diverge when there are ties in feature values or when min_samples_leaf causes skips. The true_count-based logic had the same structure, so this is an inherited bug now exposed by the refactoring. Needs audit and a test with tied feature values.
2. find_best_split: division by zero when true_mass == 0.0
In find_random_split, true_mean is guarded with if true_mass > 0f64, but in find_best_split you directly compute true_sum / true_mass without any guard. This can panic or produce NaN/inf when the first candidate split has no weighted mass on the true side.
3. Typo in doc comment
fit_with_weights in decision_tree_regressor.rs has a typo: sample_weigts (missing h).
🟡 Design & API Concerns
4. sample_weights is always f64, not generic over TX/TY
All other numeric types in the API are generic over TX: Number + FloatNumber. Hardcoding weights as &[f64] is pragmatic but inconsistent with the rest of the library's design. Reasonable simplification for a first PR, but should be tracked as a follow-up.
5. BaseTreeRegressor::fit renamed to fit_inner with silent visibility change
The original BaseTreeRegressor::fit was pub. It is now pub(crate). This is a breaking change for anyone using BaseTreeRegressor directly. The PR description notes internal methods changed but are not in the public API — this should be verified against the crate's semver surface.
6. validate_sample_weights called twice per tree in the forest path
The call chain DecisionTreeRegressor::fit_inner → BaseTreeRegressor::fit_inner → validate_sample_weights means validation is duplicated. For BaseForestRegressor, validation is called at the forest level and again in each fit_weak_learner invocation, adding O(n_trees × n_rows) overhead. Move validation to the outermost entry point only.
7. bootstrap=false + weighted path is undocumented
When bootstrap=false, weights are passed to fit_weak_learner for tree fitting but are not used for sampling. This is correct behaviour, but it is not documented, which could confuse future maintainers.
🟢 Coverage Gaps (matching Codecov report)
The Codecov report flags 26 uncovered lines, primarily in base_forest_regressor.rs (11 missing). Likely missing paths:
- The
bootstrap=false+sample_weights=Some(...)branch (weighted forest without bootstrapping) - The
WeightedIndexconstruction error path - The
balance_propertytest should be verified to exercise the non-bootstrap weighted code path explicitly
🔵 Minor / Style
sample_with_replacementhas duplicated loop bodies between theSome(dist)andNonebranches — can be collapsed into a single loop using a closure.- Blank line added at the top of the
testsmodule inbase_forest_regressor.rsis cosmetic noise. - The
balance_propertytest calls.xa(true, &x)— add a comment clarifying whatxacomputes for reviewers unfamiliar with the crate internals.
Mec-iS
left a comment
There was a problem hiding this comment.
Good implementation overall — the API design, validation logic, and test coverage are solid. Two correctness issues in find_best_split need to be fixed before merging (division-by-zero guard and the true_mass accumulation audit with tied features). The other comments are improvements rather than blockers.
Summary of blocking issues:
find_best_split: missingtrue_mass == 0.0guard (parallels the existing guard infind_random_split) — can produceNaN/infsplits silently.find_best_split:true_massaccumulation needs a test with tied feature values and non-uniform weights to confirm it is consistent withtrue_sumacross all loop paths.
Non-blocking but recommended:
- Restore the doc comment removed from
BaseForestRegressor::fit. - Fix the
sample_weigtstypo. - Move
validate_sample_weightsto outermost entry points only to avoid O(n_trees × n_rows) redundant work. - Collapse the duplicated loop bodies in
sample_with_replacement. - Add a
bootstrap=true+ weighted prediction accuracy test inbase_forest_regressor.
| } | ||
|
|
||
| /// Build a decision tree regressor from the training data. | ||
| /// * `x` - _NxM_ matrix with _N_ observations and _M_ features in each observation. |
There was a problem hiding this comment.
Typo in the doc comment: sample_weigts → sample_weights.
| @@ -435,15 +494,17 @@ impl<TX: Number + PartialOrd, TY: Number, X: Array2<TX>, Y: Array1<TY>> | |||
| { | |||
There was a problem hiding this comment.
Bug: true_mass diverges from true_sum when feature values are tied or min_samples_leaf causes a skip.
In the two continue branches above (tie-breaking and min_samples_leaf guard), true_mass and true_sum are both updated — which is correct. However, at the bottom of the loop (after the if split_score block), only true_sum and true_count are updated; true_mass is updated here but note it is not reset between candidate splits — it just keeps accumulating across the whole loop without ever being zeroed. Meanwhile true_count is also not reset, so both have the same structure. But true_mass is computed as a float product while true_count is an integer count, so any bootstrapping scenario where samples[i] > 1 will produce different values. The real concern is the combination with weighted bootstrapped samples: a row appearing multiple times will have its mass counted once per samples[i] copy correctly, but if the same row appears in both a continue path and then again at the bottom, it will be double-counted in true_mass but not in true_sum.
Please add a test with tied feature values and non-uniform weights to confirm the split scoring is correct.
| true_sum += visitor.samples[*i] as f64 * visitor.y.get(*i).to_f64().unwrap(); | ||
| true_mass += visitor.mass_of(*i); | ||
| true_sum += visitor.mass_of(*i) * visitor.y.get(*i).to_f64().unwrap(); | ||
| continue; |
There was a problem hiding this comment.
Potential division by zero / NaN: find_best_split does not guard against true_mass == 0.0.
find_random_split guards this correctly with if true_mass > 0f64 { ... } else { 0.0 }. But here true_mean = true_sum / true_mass is computed unconditionally. If all samples in the current node have zero weight and happen to be on the true side (e.g. a node created during bootstrapping where all in-bag counts are 0), this will produce NaN or inf, corrupting the gain calculation silently.
Apply the same guard as find_random_split:
let true_mean = if true_mass > 0.0 { true_sum / true_mass } else { 0.0 };
let false_mass = mass - true_mass;
let false_mean = if false_mass > 0.0 { (sum - true_sum) / false_mass } else { 0.0 };| ) -> Result<(), Failed> { | ||
| if let Some(w) = sample_weights { | ||
| if w.len() != n_rows { | ||
| return Err(Failed::fit( |
There was a problem hiding this comment.
validate_sample_weights is now called here in fit_inner, and also inside fit_weak_learner (via the same call chain) when invoked from BaseForestRegressor::fit. This means validation runs once per tree, i.e. O(n_trees × n_rows) total. Since the forest already validates at the top level before the tree loop, the per-tree call is redundant. Consider skipping validation inside fit_weak_learner — it is an internal method and the forest guarantees validity before calling it.
| pub fn fit( | ||
| x: &X, | ||
| y: &Y, | ||
| sample_weights: Option<&[f64]>, |
There was a problem hiding this comment.
The public signature of BaseForestRegressor::fit changed (new sample_weights parameter). The doc comment was removed in this hunk. Please restore it, and document the new sample_weights parameter.
| samples[xi] += 1; | ||
| } | ||
| } | ||
|
|
There was a problem hiding this comment.
The two branches of sample_with_replacement have identical loop bodies apart from the sampling call. This can be simplified to a single loop:
for _ in 0..nrows {
let xi = if let Some(dist) = distribution {
rng.sample(dist)
} else {
rng.random_range(0..nrows)
};
samples[xi] += 1;
}Or, construct a closure/enum once before the loop to avoid the per-iteration branch if performance matters.
| .sum::<f64>(); | ||
| let actual = y_hat | ||
| .iter() | ||
| .zip(normalized_weights.iter()) |
There was a problem hiding this comment.
The balance_property test covers the bootstrap=false + weighted path, which is great. However, bootstrap=true + weighted path (where WeightedIndex is actually used for sampling) is only covered by test_each_tree_gets_different_bootstrap_sample, which asserts structural properties but not prediction accuracy. Consider adding a test that uses bootstrap=true with weights and checks that the weighted predictions reflect the sampling bias (similar to fit_with_weights_predicts_approx_weighted_mean in random_forest_regressor).
|
@slievens thanks, please see the review items |
Following changes were made: - Guard against true_mass and false_mass being zero. Test added. - Validate sample_weights only once per call to fit. Directly in fit_with_weights. - Fix typo in documentation. - mplementation of sample_with_replacement is nicer. - Removed documentation was restored. - Add a test with tied feature values and different sample_weights.
Add tests on base_forest_regressor and extra_trees_regressor to check that the root does predict the weighted mean when sample weights are used.
Mec-iS
left a comment
There was a problem hiding this comment.
All blocking issues from the previous review have been resolved — great follow-up.
✅ true_mass / true_sum divergence: fixed and covered by the new weights_on_tied_feature_values test.
✅ Division-by-zero in find_best_split: guarded symmetrically with find_random_split, with two dedicated zero-weight tests.
✅ Validation moved to the outermost entry points; no more O(n_trees × n_rows) overhead.
✅ Typo fixed, doc comment restored, sample_with_replacement loop collapsed, bootstrap=true weighted coverage added.
Three minor follow-up notes left inline (one perf nit on mass recomputation in split, one wrong comment, one design note on pub(crate) visibility). None of these block merging.
| }; | ||
| let sum = self.nodes()[visitor.node].output * mass; | ||
|
|
||
| let mut variables = (0..n_attr).collect::<Vec<_>>(); |
There was a problem hiding this comment.
Minor: recomputes mass every time split is entered, even when weights are absent.
let mass = match visitor.sample_weights {
Some(_) => (0..visitor.samples.len()).map(|i| visitor.mass_of(i)).sum(),
None => n as f64,
};When sample_weights is Some, this is an O(n) loop on every split call — once per (node × attribute). The total is O(n × n_nodes × mtry). For moderate datasets this is fine, but it could be avoided cheaply: mass is already computed in fit_weak_learner and passed down to both splitters as a parameter. Consider passing mass as an argument to split the same way it is already passed to find_best_split and find_random_split, and remove this re-computation block entirely.
| } | ||
| } | ||
| } | ||
|
|
There was a problem hiding this comment.
Nit: the comment in the test says // No bootstrapping but bootstrap: true is set. Should read // Use bootstrapping.
|
|
||
| /// Validates the sample weights | ||
| pub(crate) fn validate_sample_weights(sample_weights: &[f64], n_rows: usize) -> Result<(), Failed> { | ||
| if sample_weights.len() != n_rows { |
There was a problem hiding this comment.
Design note (non-blocking): validate_sample_weights is now pub(crate), which means every public entry point (DecisionTreeRegressor::fit_with_weights, RandomForestRegressor::fit_with_weights, ExtraTreesRegressor::fit_with_weights) re-exports it indirectly. This is fine for the current codebase, but it would be cleaner to keep this function private to the module and duplicate the three small validation call-sites rather than leaking an internal helper. Not a blocker, just worth a note for future API hygiene.
Add
fit_with_weightsto DecisionTreeRegressor, RandomForestRegressor and ExtraTreesRegressor.fitkeeps its signature and delegates with no weights.sample_weightsinfit. Bootstrap draws useWeightedIndexwhen weights are given. Pass weights to each tree.Checklist
Current behaviour
Sample weights are not supported on tree based regressors.
New expected behaviour
Sample weights are supported on tree based regressors.
Change logs
Added
New method on public API "fit_with_weights"
Changed
Some internal methods have a slightly different name and signature. These are not in the public API.