Skip to content

feat(tree,ensemble): add sample weights to regressor fit - #467

Merged
Mec-iS merged 3 commits into
smartcorelib:mainfrom
slievens:rf-sample-weights-on-fit-v2
Sep 27, 2026
Merged

Mec-iS merged 3 commits into
smartcorelib:mainfrom
slievens:rf-sample-weights-on-fit-v2

Conversation

@slievens

Copy link
Copy Markdown
Contributor

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

Checklist

  • [ x] My branch is up-to-date with main branch.
  • [ x] Everything works and tested on latest stable Rust.
  • [ x] Coverage and Linting have been applied

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.

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

codecov Bot commented Sep 26, 2026 •

Copy link
Copy Markdown

Codecov Report

❌ Patch coverage is 70.21277% with 28 lines in your changes missing coverage. Please review.
✅ Project coverage is 63.88%. Comparing base (9eaae9e) to head (e0202e9).
⚠️ Report is 189 commits behind head on main.

Files with missing lines Patch % Lines
src/tree/base_tree_regressor.rs 76.66% 14 Missing ⚠️
src/ensemble/base_forest_regressor.rs 44.44% 10 Missing ⚠️
src/ensemble/extra_trees_regressor.rs 66.66% 2 Missing ⚠️
src/ensemble/random_forest_regressor.rs 80.00% 1 Missing ⚠️
src/tree/decision_tree_regressor.rs 80.00% 1 Missing ⚠️
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.
📢 Have feedback on the report? Share it here.

🚀 New features to boost your workflow:
  • ❄️ Test Analytics: Detect flaky tests, report on failures, and find test suite problems.

@Mec-iS Mec-iS left a comment

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

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 WeightedIndex construction error path
  • The balance_property test should be verified to exercise the non-bootstrap weighted code path explicitly

🔵 Minor / Style

  • sample_with_replacement has duplicated loop bodies between the Some(dist) and None branches — can be collapsed into a single loop using a closure.
  • Blank line added at the top of the tests module in base_forest_regressor.rs is cosmetic noise.
  • The balance_property test calls .xa(true, &x) — add a comment clarifying what xa computes for reviewers unfamiliar with the crate internals.

@Mec-iS Mec-iS left a comment

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

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:

  1. find_best_split: missing true_mass == 0.0 guard (parallels the existing guard in find_random_split) — can produce NaN/inf splits silently.
  2. find_best_split: true_mass accumulation needs a test with tied feature values and non-uniform weights to confirm it is consistent with true_sum across all loop paths.

Non-blocking but recommended:

  • Restore the doc comment removed from BaseForestRegressor::fit.
  • Fix the sample_weigts typo.
  • Move validate_sample_weights to 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 in base_forest_regressor.

}

/// Build a decision tree regressor from the training data.
/// * `x` - _NxM_ matrix with _N_ observations and _M_ features in each observation.

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

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>>
{

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

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;

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

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 };

Comment thread src/tree/base_tree_regressor.rs Outdated
) -> Result<(), Failed> {
if let Some(w) = sample_weights {
if w.len() != n_rows {
return Err(Failed::fit(

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

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]>,

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

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;
}
}

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

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())

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

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).

@Mec-iS

Mec-iS commented Sep 26, 2026

Copy link
Copy Markdown
Collaborator

@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.
@slievens

Copy link
Copy Markdown
Contributor Author

@Mec-iS I hope the additional two commits (0ed5b0a and e0202e9) resolve the items in your commit.

@Mec-iS Mec-iS left a comment

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

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<_>>();

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

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.

}
}
}

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

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 {

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

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.

@Mec-iS Mec-iS changed the title WIP feat(tree,ensemble): add sample weights to regressor fit feat(tree,ensemble): add sample weights to regressor fit Sep 27, 2026
@Mec-iS
Mec-iS merged commit cbbb32a into smartcorelib:main Sep 27, 2026
15 checks passed
Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Labels

None yet

Projects

None yet

Development

Successfully merging this pull request may close these issues.

2 participants