Skip to content

Add BaseModel.set_pos_weight for imbalanced binary/multilabel outcomes - #1267

Merged
jhnwu3 merged 1 commit into
sunlabuiuc:masterfrom
solarsys:feat/pos-weight
Oct 8, 2026
Merged

jhnwu3 merged 1 commit into
sunlabuiuc:masterfrom
solarsys:feat/pos-weight

Conversation

@solarsys

@solarsys solarsys commented Oct 3, 2026

Copy link
Copy Markdown
Collaborator

Problem

For binary and multilabel labels, BaseModel.get_loss_function() returns plain F.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

model = RNN(dataset=samples)
model.set_pos_weight("balanced", train_dataset)   # negatives / positives, per label
model.set_pos_weight(9.0)                         # or a fixed weight
model.set_pos_weight(None)                        # back to unweighted
model.pos_weight                                  # tensor or None
  • get_loss_function() passes the weight as pos_weight, moved to the logits' device. The default is unchanged.
  • Multilabel: "balanced" gives one weight per label; a label with no positives gets weight 1 and a warning.
  • Multiclass and regression raise ValueError.

Design choices, open to your view

  • A method, not a constructor argument (the spec proposed BaseModel(..., pos_weight=…)). Of the 29 models, none pass extra arguments through to BaseModel.__init__, so a constructor argument couldn't reach them without editing each one. Every model fetches the loss with self.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.dataset is usually the full, unsplit sample set, and computing prevalence from it would let validation and test labels set the weight.
  • Not covered: GAMENet, SafeDrug, MoleRec and MICRON compute their own losses. The docs say so.

Tests, docs, example

  • New tests/core/test_pos_weight.py (8 tests):
    • the default is unweighted;
    • a fixed weight matches F.binary_cross_entropy_with_logits(pos_weight=…);
    • "balanced" uses the given training data, not the model's own;
    • "balanced" without a dataset raises;
    • clearing the weight works;
    • the forward-pass loss uses the weight;
    • multilabel gets per-label weights;
    • multiclass is rejected.
  • Existing model.mode, RNN and MLP tests pass.
  • docs/api/models.rst gets an "Imbalanced outcomes" section: pass the training split, which models ignore the weight, and that weighting shifts probabilities upward, so check calibration.
  • New examples/imbalanced_outcome_pos_weight.py: a 10% outcome, trained with and without "balanced", scored side by side.
  • Full core suite: Ran 1405 tests … OK (skipped=76). tools/check_pr_rules.py passes.

🤖 Generated with Claude Code

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

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.

@jhnwu3
jhnwu3 merged commit eb55f9c into sunlabuiuc:master Oct 8, 2026
2 checks passed
@solarsys
solarsys deleted the feat/pos-weight branch October 9, 2026 04:30
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