Skip to content

fix: strip embedded penalty terms from causal_loss_metric_gen (v0.1.1) - #21

Merged
gmgeorg merged 1 commit into
mainfrom
fix/2
Aug 23, 2026
Merged

gmgeorg merged 1 commit into
mainfrom
fix/2

Conversation

@gmgeorg

@gmgeorg gmgeorg commented Aug 23, 2026 •

Copy link
Copy Markdown
Owner

Summary

  • causal_loss_metric_gen reused the exact outcome_loss/treatment_loss instances it was given, so any penalty folded into either loss's own call() (e.g. a within-state balance penalty) would leak into the metric used for EarlyStopping/ReduceLROnPlateau/checkpoint selection and any hyperparameter-tuning objective built on it.
  • Adds an explicit penalty_free() contract to OutcomeLoss/TreatmentLoss: each reconstructs itself from a hand-written list of its own known-safe constructor args (never by scanning attribute names). A future penalty either defaults to "off" automatically or raises loudly until the subclass explicitly overrides penalty_free() — it can never leak through unnoticed. causal_loss_metric_gen now calls it on both losses before wrapping them.
  • Bumps version to 0.1.1, adds a CHANGELOG.md entry, and a dev-docs/ writeup with the full rationale (local/gitignored, per repo convention).

Test plan

  • pytest pypsps/tests/test_losses.py pypsps/tests/test_models.py pypsps/tests/test_metrics.py — 26 passed
  • pytest pypsps/tests/ (full suite) — 67 passed
  • New regression tests: penalty-free identity when no penalty declared, penalty zeroed without mutating the live loss, fail-loud for an undeclared required penalty arg, and an end-to-end check that causal_loss_metric_gen's output matches the unpenalized likelihood, not the penalized one.

causal_loss_metric_gen reused the exact outcome_loss/treatment_loss
instances it was given, so any penalty folded into either loss's own
call() (e.g. a within-state balance penalty) would leak into the
metric used for EarlyStopping/ReduceLROnPlateau/checkpoint selection.
Adds an explicit penalty_free() contract to OutcomeLoss/TreatmentLoss
that reconstructs each loss from its known-safe constructor args
(never by attribute-name sniffing), and wires it into
causal_loss_metric_gen.

Co-Authored-By: Claude Sonnet 5 <noreply@anthropic.com>
Claude-Session: https://claude.ai/code/session_015QCjbByNYkymYjyzyLaiv9
@gmgeorg
gmgeorg merged commit 13febd5 into main Aug 23, 2026
1 check 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.

1 participant