diff --git a/CHANGELOG.md b/CHANGELOG.md index dda8244..000c857 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -5,6 +5,27 @@ All notable changes to `pypsps` will be documented here. The format is based on [Keep a Changelog](https://keepachangelog.com/en/1.0.0/), and this project adheres to [Semantic Versioning](https://semver.org/spec/v2.0.0.html). +## pypsps v0.1.1 - Aug 22, 2026 + +### Fixed + +* `causal_loss_metric_gen` (`pypsps/keras/metrics.py`) reused the exact `outcome_loss`/ + `treatment_loss` instances it was given, so any penalty term 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 -- silently turning "pure held-out + likelihood" into "likelihood plus whatever penalty this candidate happened to draw", + systematically rewarding weaker penalties regardless of fit quality. No such penalty is + wired into committed `OutcomeLoss`/`TreatmentLoss` yet, so this was a structural gap rather + than an active bug, flagged as a pending follow-up in + [dev-docs/20260820-bugfixes-v1.md item 5](dev-docs/20260820-bugfixes-v1.md). Fixed by adding + an explicit `penalty_free()` method to `OutcomeLoss`/`TreatmentLoss` that reconstructs the + loss from a hand-written list of its own non-penalty constructor arguments (not by scanning + attribute names for a naming convention); `causal_loss_metric_gen` now always calls it before + wrapping the losses in its internal `CausalLoss`. A future subclass that adds a penalty either + gets it correctly zeroed automatically (if it defaults to "off") or `penalty_free()` raises + until that subclass explicitly overrides it -- it can never leak through unnoticed. See + [dev-docs/20260822-bugfixes-v3.md](dev-docs/20260822-bugfixes-v3.md) for the full writeup. + ## pypsps v0.1.0 - Aug 21, 2026 Breaking: the propensity head's output shape changed (state-conditional instead of diff --git a/pyproject.toml b/pyproject.toml index 0bd812d..1887b2e 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -1,6 +1,6 @@ [tool.poetry] name = "pypsps" -version = "0.1.0" +version = "0.1.1" description = "Predictive State Propensity Subclassification (PSPS) in Python (keras)" authors = ["Georg M. Goerg "] license = "MIT" diff --git a/pypsps/keras/losses.py b/pypsps/keras/losses.py index 4b2f794..5ad0d34 100644 --- a/pypsps/keras/losses.py +++ b/pypsps/keras/losses.py @@ -115,6 +115,28 @@ def __init__( self._n_outcome_pred_cols = n_outcome_pred_cols self._n_treatment_pred_cols = n_treatment_pred_cols + def penalty_free(self) -> "OutcomeLoss": + """Returns a fresh instance of this exact class with no penalty terms applied. + + Rebuilds `type(self)(...)` from only the arguments listed below -- never by copying + `self.__dict__` or scanning for attributes by naming convention. This is the explicit + contract for stripping penalties out of a metric like `causal_loss_metric_gen`: if a + subclass's `call()` ever adds a penalty (e.g. a within-state balance term), that + subclass MUST override this method to omit (or zero) that penalty's constructor + argument here. Forgetting to do so either raises (if the new argument is required, + since it's missing from this call) or silently keeps the penalty out (if it defaults + to "off"), but never lets an unrecognized penalty leak through. + """ + return type(self)( + loss=self._loss, + treatment_loss=self._treatment_loss, + n_outcome_true_cols=self._n_outcome_true_cols, + n_outcome_pred_cols=self._n_outcome_pred_cols, + n_treatment_pred_cols=self._n_treatment_pred_cols, + reduction=self.reduction, + name=self.name, + ) + def call(self, y_true, y_pred): """Evaluates Causal Loss on (y_true, y_pred) for binary loss and Normal outcomes. @@ -224,6 +246,23 @@ def __init__( self._n_outcome_pred_cols = n_outcome_pred_cols self._n_treatment_pred_cols = n_treatment_pred_cols + def penalty_free(self) -> "TreatmentLoss": + """Returns a fresh instance of this exact class with no penalty terms applied. + + See `OutcomeLoss.penalty_free` for the contract: rebuilds `type(self)(...)` from only + the arguments below, never by copying `self.__dict__` or scanning attribute names. A + subclass whose `call()` adds a penalty (e.g. a within-state balance term) MUST + override this method to omit or zero that penalty's constructor argument. + """ + return type(self)( + loss=self._loss, + n_outcome_true_cols=self._n_outcome_true_cols, + n_outcome_pred_cols=self._n_outcome_pred_cols, + n_treatment_pred_cols=self._n_treatment_pred_cols, + reduction=self.reduction, + name=self.name, + ) + def call(self, y_true, y_pred): """Evaluates the marginal treatment (dose) loss -log p(a | x).""" n_states = utils.get_n_states( diff --git a/pypsps/keras/metrics.py b/pypsps/keras/metrics.py index 6f97269..0635baa 100644 --- a/pypsps/keras/metrics.py +++ b/pypsps/keras/metrics.py @@ -186,6 +186,20 @@ def predictive_state_df(y_true, y_pred) -> tf.Tensor: return predictive_state_df +def _penalty_free_copy(loss_obj): + """Returns a penalty-free instance of `loss_obj`, without mutating `loss_obj` itself + (which may still be the live training-time loss). + + Delegates to `loss_obj.penalty_free()` (see `losses.OutcomeLoss.penalty_free` / + `losses.TreatmentLoss.penalty_free`), which rebuilds the object from an explicit, + class-declared list of constructor arguments -- never by scanning attribute names for a + naming convention like `_lambda_*`. That keeps this correct even for a penalty that isn't + named with a `lambda`-style prefix: the loss class itself, not this function, decides + what counts as a penalty. + """ + return loss_obj.penalty_free() + + def causal_loss_metric_gen( outcome_loss: losses.OutcomeLoss, treatment_loss: losses.TreatmentLoss, @@ -200,6 +214,12 @@ def causal_loss_metric_gen( causal_loss = outcome_loss_weight * outcome_loss(y_true, y_pred) + alpha * treatment_loss(y_true, y_pred) + with `outcome_loss` and `treatment_loss` stripped of any embedded penalty terms (see + `_penalty_free_copy`) and wrapped in a `CausalLoss` with `predictive_states_regularizer= + None`, so this metric always reports the exact joint likelihood, never a penalized/ + regularized proxy for it -- regardless of what alpha, outcome_loss_weight, or predictive + state / balance penalties the model was actually trained with. + This metric function can be passed to model.compile(metrics=[...]). Parameters @@ -218,10 +238,12 @@ def causal_loss_metric_gen( function A function metric that takes (y_true, y_pred) and returns the causal loss as a float value (can be passed as metric). """ - # Construct an instance of CausalLoss with the given parameters. + # Construct an instance of CausalLoss with the given parameters, using penalty-free + # copies of outcome_loss/treatment_loss so this metric can never inherit a penalty term + # baked into either of them. causal_loss_obj = losses.CausalLoss( - outcome_loss=outcome_loss, - treatment_loss=treatment_loss, + outcome_loss=_penalty_free_copy(outcome_loss), + treatment_loss=_penalty_free_copy(treatment_loss), alpha=alpha, outcome_loss_weight=outcome_loss_weight, ) diff --git a/pypsps/tests/test_metrics.py b/pypsps/tests/test_metrics.py index 61185fc..62ebe34 100644 --- a/pypsps/tests/test_metrics.py +++ b/pypsps/tests/test_metrics.py @@ -1,7 +1,7 @@ import numpy as np import tensorflow as tf -from pypsps.keras import metrics +from pypsps.keras import losses, metrics, neglogliks def _make_y_pred(outcome_blocks, weights, treatment_blocks): @@ -200,3 +200,149 @@ def test_predictive_state_df_gen(): result = func(None, y_pred) # Check that result is a scalar tensor. assert result.shape.ndims == 0 or (result.shape.ndims == 1 and result.shape[0] == 1) + + +def _make_treatment_loss(**overrides): + """Builds a plain (no-penalty) `TreatmentLoss` with sensible test defaults.""" + kwargs = dict( + loss=tf.keras.losses.BinaryCrossentropy(reduction="none"), + n_outcome_true_cols=1, + n_outcome_pred_cols=2, + n_treatment_pred_cols=1, + reduction="sum_over_batch_size", + ) + kwargs.update(overrides) + return losses.TreatmentLoss(**kwargs) + + +def test_treatment_loss_penalty_free_is_identity_when_no_penalty_declared(): + """Base `TreatmentLoss.penalty_free()` has no penalty args to drop, so it must return an + equivalent (but distinct) instance -- not the same object.""" + original = _make_treatment_loss() + clean = original.penalty_free() + + assert clean is not original + assert clean._n_outcome_true_cols == original._n_outcome_true_cols + assert clean._n_outcome_pred_cols == original._n_outcome_pred_cols + assert clean._n_treatment_pred_cols == original._n_treatment_pred_cols + assert clean._loss is original._loss + + +class _BalancePenalizedTreatmentLoss(losses.TreatmentLoss): + """Test double simulating a hypothetical within-state balance penalty folded directly + into `TreatmentLoss.call` (`self._lambda_balance * balance_penalty`), the pattern + `causal_loss_metric_gen` must be robust to -- without relying on the `lambda_balance` + name, since `penalty_free()` is reconstructed via explicit constructor arguments, not + attribute-name sniffing.""" + + def __init__(self, *args, lambda_balance: float = 0.0, **kwargs): + """Stores `lambda_balance`, the (test-only) penalty weight.""" + super().__init__(*args, **kwargs) + self._lambda_balance = lambda_balance + + def call(self, y_true, y_pred): + """Adds a constant `lambda_balance` penalty on top of the real treatment NLL.""" + return super().call(y_true, y_pred) + self._lambda_balance + + +def test_treatment_loss_penalty_free_zeros_declared_penalty_without_mutating_original(): + """A subclass that inherits `penalty_free()` without overriding it drops any constructor + argument not explicitly forwarded by the base implementation -- here, `lambda_balance` + falls back to its own "off" default (0.0) -- and the original, live instance must be + left untouched.""" + original = _BalancePenalizedTreatmentLoss( + loss=tf.keras.losses.BinaryCrossentropy(reduction="none"), + n_outcome_true_cols=1, + n_outcome_pred_cols=2, + n_treatment_pred_cols=1, + lambda_balance=5.0, + reduction="sum_over_batch_size", + ) + clean = original.penalty_free() + + assert clean is not original + assert clean._lambda_balance == 0.0 + assert original._lambda_balance == 5.0 + + +def test_treatment_loss_penalty_free_fails_loudly_for_undeclared_required_penalty_arg(): + """If a subclass adds a *required* penalty argument (no safe "off" default) and doesn't + override `penalty_free()` to account for it, reconstruction must raise rather than + silently guess a value -- forcing whoever adds the penalty to explicitly decide how + `penalty_free()` should handle it.""" + + class _RequiredPenaltyTreatmentLoss(losses.TreatmentLoss): + """Test double whose penalty weight has no safe "off" default.""" + + def __init__(self, *args, lambda_balance: float, **kwargs): + """Stores the required `lambda_balance` penalty weight.""" + super().__init__(*args, **kwargs) + self._lambda_balance = lambda_balance + + original = _RequiredPenaltyTreatmentLoss( + loss=tf.keras.losses.BinaryCrossentropy(reduction="none"), + n_outcome_true_cols=1, + n_outcome_pred_cols=2, + n_treatment_pred_cols=1, + lambda_balance=5.0, + reduction="sum_over_batch_size", + ) + try: + original.penalty_free() + assert False, "expected TypeError: lambda_balance is required and not forwarded" + except TypeError: + pass + + +def test_causal_loss_metric_gen_strips_embedded_treatment_loss_penalty(): + """causal_loss_metric_gen must report the exact joint likelihood even when the passed-in + treatment_loss instance carries an embedded penalty term -- the metric used for + EarlyStopping/checkpoint selection/Optuna must never be a penalized proxy.""" + n_outcome_pred_cols, n_treatment_pred_cols, n_outcome_true_cols = 2, 1, 1 + weights = [[0.5, 0.5], [0.3, 0.7], [0.9, 0.1]] + y_pred = _make_y_pred( + outcome_blocks=[np.zeros((3, 2)), np.ones((3, 2))], + weights=weights, + treatment_blocks=[[[0.9, 0.7], [0.5, 0.5], [0.2, 0.6]]], + ) + y_true = tf.constant([[5.0, 1.0], [3.0, 0.0], [8.0, 1.0]], dtype=tf.float32) + + def _build(lambda_balance): + """Builds a matching (outcome_loss, treatment_loss) pair for a given penalty weight.""" + outcome_loss = losses.OutcomeLoss( + loss=neglogliks.NegloglikNormal(reduction="none"), + treatment_loss=tf.keras.losses.BinaryCrossentropy(reduction="none"), + n_outcome_true_cols=n_outcome_true_cols, + n_outcome_pred_cols=n_outcome_pred_cols, + n_treatment_pred_cols=n_treatment_pred_cols, + reduction="sum_over_batch_size", + ) + treatment_loss = _BalancePenalizedTreatmentLoss( + loss=tf.keras.losses.BinaryCrossentropy(reduction="none"), + n_outcome_true_cols=n_outcome_true_cols, + n_outcome_pred_cols=n_outcome_pred_cols, + n_treatment_pred_cols=n_treatment_pred_cols, + lambda_balance=lambda_balance, + reduction="sum_over_batch_size", + ) + return outcome_loss, treatment_loss + + penalized_outcome_loss, penalized_treatment_loss = _build(lambda_balance=1000.0) + clean_outcome_loss, clean_treatment_loss = _build(lambda_balance=0.0) + + metric_fn = metrics.causal_loss_metric_gen( + outcome_loss=penalized_outcome_loss, treatment_loss=penalized_treatment_loss + ) + metric_value = metric_fn(y_true, y_pred).numpy() + + expected_clean = ( + clean_outcome_loss(y_true, y_pred) + clean_treatment_loss(y_true, y_pred) + ).numpy() + expected_penalized = ( + penalized_outcome_loss(y_true, y_pred) + penalized_treatment_loss(y_true, y_pred) + ).numpy() + + np.testing.assert_allclose(metric_value, expected_clean, rtol=1e-5) + assert not np.isclose(metric_value, expected_penalized, rtol=1e-5) + # The original, live training-time loss must be unmodified. + assert penalized_treatment_loss._lambda_balance == 1000.0