Repository navigation
Add BaseModel.set_pos_weight for imbalanced binary/multilabel outcomes - #1267
Merged
Merged
Conversation
The default loss for binary and multilabel labels was plain binary_cross_entropy_with_logits with no way to weight the positives, which rare clinical outcomes (5-15% prevalence) usually need. - BaseModel.set_pos_weight(pos_weight, dataset=None): a number, one number per label, "balanced" (negatives / positives per label in `dataset`), or None to remove it. "balanced" requires the dataset explicitly, so the weight comes from the training split rather than the model's own, usually unsplit, dataset. Multiclass/regression raise ValueError. - get_loss_function passes the weight as pos_weight (moved to the logits' device). Models fetch the loss per forward pass, so a weight set after construction applies to every model using the default loss. Default behaviour is unchanged. - tests/core/test_pos_weight.py: unweighted default, fixed weight, balanced from the given training data (not the model's data), errors, clearing, the forward loss, per-label multilabel weights, multiclass. - docs/api/models.rst: "Imbalanced outcomes", including which models compute their own loss and the effect on calibration. - examples/imbalanced_outcome_pos_weight.py. Co-Authored-By: Claude Opus 5.5 <[email protected]>
jhnwu3
approved these changes
Oct 8, 2026
jhnwu3
left a comment
Collaborator
There was a problem hiding this comment.
lgtm.
This might be worth revisiting in a later timeline to see if there's a cleaner solution to this problem with the outdated pyhealth.trainer. But, I think this might be the easiest for now.
This file contains hidden or bidirectional Unicode text that may be interpreted or compiled differently than what appears below. To review, open the file in an editor that reveals hidden Unicode characters.
Learn more about bidirectional Unicode characters
Sign up for free
to join this conversation on GitHub.
Already have an account?
Sign in to comment
Add this suggestion to a batch that can be applied as a single commit.This suggestion is invalid because no changes were made to the code.Suggestions cannot be applied while the pull request is closed.Suggestions cannot be applied while viewing a subset of changes.Only one suggestion per line can be applied in a batch.Add this suggestion to a batch that can be applied as a single commit.Applying suggestions on deleted lines is not supported.You must change the existing code in this line in order to create a valid suggestion.Outdated suggestions cannot be applied.This suggestion has been applied or marked resolved.Suggestions cannot be applied from pending reviews.Suggestions cannot be applied on multi-line comments.Suggestions cannot be applied while the pull request is queued to merge.Suggestion cannot be applied right now. Please check back later.
Problem
For binary and multilabel labels,
BaseModel.get_loss_function()returns plainF.binary_cross_entropy_with_logits, and there is no way to weight the positive class. Rare clinical outcomes (5–15% prevalence is typical) usually need that. A downstream EHR project had to patch around it.Change
get_loss_function()passes the weight aspos_weight, moved to the logits' device. The default is unchanged."balanced"gives one weight per label; a label with no positives gets weight 1 and a warning.ValueError.Design choices, open to your view
BaseModel(..., pos_weight=…)). Of the 29 models, none pass extra arguments through toBaseModel.__init__, so a constructor argument couldn't reach them without editing each one. Every model fetches the loss withself.get_loss_function()on each forward pass (31 call sites), so a weight set after construction applies to all of them."balanced"needs the dataset passed explicitly.model.datasetis usually the full, unsplit sample set, and computing prevalence from it would let validation and test labels set the weight.Tests, docs, example
tests/core/test_pos_weight.py(8 tests):F.binary_cross_entropy_with_logits(pos_weight=…);"balanced"uses the given training data, not the model's own;"balanced"without a dataset raises;model.mode, RNN and MLP tests pass.docs/api/models.rstgets an "Imbalanced outcomes" section: pass the training split, which models ignore the weight, and that weighting shifts probabilities upward, so check calibration.examples/imbalanced_outcome_pos_weight.py: a 10% outcome, trained with and without"balanced", scored side by side.Ran 1405 tests … OK (skipped=76).tools/check_pr_rules.pypasses.🤖 Generated with Claude Code