Conversation
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
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.
Summary
causal_loss_metric_genreused the exactoutcome_loss/treatment_lossinstances it was given, so any penalty folded into either loss's owncall()(e.g. a within-state balance penalty) would leak into the metric used forEarlyStopping/ReduceLROnPlateau/checkpoint selection and any hyperparameter-tuning objective built on it.penalty_free()contract toOutcomeLoss/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 overridespenalty_free()— it can never leak through unnoticed.causal_loss_metric_gennow calls it on both losses before wrapping them.0.1.1, adds aCHANGELOG.mdentry, and adev-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 passedpytest pypsps/tests/(full suite) — 67 passedcausal_loss_metric_gen's output matches the unpenalized likelihood, not the penalized one.