Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension


Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
21 changes: 21 additions & 0 deletions CHANGELOG.md
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down
2 changes: 1 addition & 1 deletion pyproject.toml
Original file line number Diff line number Diff line change
@@ -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 <im@gmge.org>"]
license = "MIT"
Expand Down
39 changes: 39 additions & 0 deletions pypsps/keras/losses.py
Original file line number Diff line number Diff line change
Expand Up @@ -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.

Expand Down Expand Up @@ -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(
Expand Down
28 changes: 25 additions & 3 deletions pypsps/keras/metrics.py
Original file line number Diff line number Diff line change
Expand Up @@ -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,
Expand All @@ -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
Expand All @@ -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,
)
Expand Down
148 changes: 147 additions & 1 deletion pypsps/tests/test_metrics.py
Original file line number Diff line number Diff line change
@@ -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):
Expand Down Expand Up @@ -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
Loading