Skip to content

fix: state-conditional propensity head + posterior-mixing correctness (v0.1.0) - #20

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

gmgeorg merged 1 commit into
mainfrom
fix/1

Conversation

@gmgeorg

@gmgeorg gmgeorg commented Aug 22, 2026 •

Copy link
Copy Markdown
Owner

Squashes the fix/1 branch's 8 commits into one for review. Summary of what changed and why:

Core correctness fix

  • OutcomeLoss mixed predictive states using the prior weights P(s_k|x) instead of the posterior P(s_k|x,a); since treatment a is observed, only the posterior makes TreatmentLoss + OutcomeLoss telescope exactly into -log p(a,y|x). Computing the posterior requires each state's own P(a|s_k,x), so the propensity head is now state-conditional (one BiasOnly+sigmoid per state, like the outcome head) instead of pre-mixed into a single marginal column -- a breaking change to the model's output shape; models trained/saved with <0.1.0 are not compatible with this release.
  • log(weights + eps) in the state-mixture losses clamped log-weights at log(eps), distorting sharp/near-degenerate predictive states; replaced with a safe_log that only substitutes for exact-zero entries.
  • CausalLoss's alpha (treatment-loss weight) was a plain Python float. model.fit() traces CausalLoss.call() into a compiled graph (run_eagerly=False by default), so the float got baked in as a constant at the first trace; AlphaScheduleCallback rebinding loss._alpha between epochs only updated the eager Python attribute, which the compiled training function never read again -- the annealing schedule silently never reached the optimizer, despite reporting the intended curve. alpha is now a tf.Variable, updated via .assign(...). Verified directly (tf.function trace before/after) and end-to-end via model.fit().
  • predict_ute_binary/predict_ute_continuous hardcoded n_outcome_pred_cols=2 (Normal-only), silently mis-slicing columns -- and computing the wrong number of states -- for any other outcome distribution (e.g. the exponential/survival model), returning a plausible-looking but meaningless number with no error. Now converts each outcome distribution's raw parameters to its actual mean (Normal: loc; Exponential: exp(-log_rate)), raising NotImplementedError for any distribution it doesn't know how to convert.
  • build_toy_model/build_model_binary_normal didn't compile causal_loss_metric (only build_model_binary_exponential did), so get_default_callbacks' documented recommended monitor ("val_causal_loss_metric") didn't exist for them; Keras only warns, never raises, on a missing monitored metric, so EarlyStopping/ReduceLROnPlateau silently never triggered.

New helpers

  • utils.get_column_layout(model) / utils.ColumnLayout: reads n_outcome_pred_cols/n_treatment_pred_cols/n_outcome_true_cols off a compiled model's CausalLoss instead of hardcoding them at each call site -- the root cause of a real shape-mismatch bug found in a demo notebook, and the same pattern that caused the predict_ute_* bug above.
  • models.get_propensity_state_conditional_means(model): since the propensity head is no longer a single pypress.keras.layers.PredictiveStateMeans layer, reconstructs the same per-state constants (post-sigmoid) from the model's propensity_logit_state_ layers -- the replacement for the old model.layers[-2].state_conditional_means sanity check.
  • metrics.uniform_entropy_gen: mean Shannon entropy of predictive state weights across a batch.

Other fixes

  • BiasOnly.get_config() dropped units (silently reverting to 1 on reload) and stored bias_regularizer raw instead of via tf.keras.regularizers.serialize/deserialize.
  • The three bootstrap_* functions seeded their resampling RNG inconsistently (RandomState(0) vs RandomState(n_samples)); all three now take an explicit random_state: int = 0.
  • README.md's and docs/inference.md's code examples called a nonexistent inference.predict_ate(...) and called utils.split_y_pred(preds) with the wrong arity/unpacking -- anyone copy-pasting the README example hit an immediate error. Fixed and verified end-to-end.
  • recommended_callbacks renamed to get_default_callbacks; monitor is now a required argument instead of defaulting to "val_loss".
  • Bumped pypress dependency v0.0.6 -> v0.2.3; Uniform regularizer's l1 penalty argument renamed to l2 to track pypress's L1 -> L2 entropy penalty change.
  • Added CHANGELOG.md, backfilling v0.0.1-v0.0.13 and documenting v0.1.0.

Version bumped to 0.1.0 in pyproject.toml (breaking propensity-head shape change). Re-ran and fixed the tracked demo/example notebooks against the new API

… (v0.1.0)

Squashes the fix/1 branch's 8 commits into one for review. Summary of what
changed and why:

Core correctness fix
--------------------
* OutcomeLoss mixed predictive states using the prior weights P(s_k|x) instead
  of the posterior P(s_k|x,a); since treatment `a` is observed, only the
  posterior makes TreatmentLoss + OutcomeLoss telescope exactly into
  -log p(a,y|x). Computing the posterior requires each state's own
  P(a|s_k,x), so the propensity head is now state-conditional (one
  BiasOnly+sigmoid per state, like the outcome head) instead of pre-mixed into
  a single marginal column -- a breaking change to the model's output shape;
  models trained/saved with <0.1.0 are not compatible with this release.
* `log(weights + eps)` in the state-mixture losses clamped log-weights at
  log(eps), distorting sharp/near-degenerate predictive states; replaced with
  a safe_log that only substitutes for exact-zero entries.
* CausalLoss's `alpha` (treatment-loss weight) was a plain Python float.
  model.fit() traces CausalLoss.call() into a compiled graph
  (run_eagerly=False by default), so the float got baked in as a constant at
  the first trace; AlphaScheduleCallback rebinding `loss._alpha` between
  epochs only updated the eager Python attribute, which the compiled training
  function never read again -- the annealing schedule silently never reached
  the optimizer, despite reporting the intended curve. `alpha` is now a
  tf.Variable, updated via `.assign(...)`. Verified directly (tf.function
  trace before/after) and end-to-end via model.fit().
* predict_ute_binary/predict_ute_continuous hardcoded n_outcome_pred_cols=2
  (Normal-only), silently mis-slicing columns -- and computing the wrong
  number of states -- for any other outcome distribution (e.g. the
  exponential/survival model), returning a plausible-looking but meaningless
  number with no error. Now converts each outcome distribution's raw
  parameters to its actual mean (Normal: loc; Exponential: exp(-log_rate)),
  raising NotImplementedError for any distribution it doesn't know how to
  convert.
* build_toy_model/build_model_binary_normal didn't compile
  `causal_loss_metric` (only build_model_binary_exponential did), so
  get_default_callbacks' documented recommended monitor
  ("val_causal_loss_metric") didn't exist for them; Keras only warns, never
  raises, on a missing monitored metric, so EarlyStopping/ReduceLROnPlateau
  silently never triggered.

New helpers
-----------
* utils.get_column_layout(model) / utils.ColumnLayout: reads
  n_outcome_pred_cols/n_treatment_pred_cols/n_outcome_true_cols off a
  compiled model's CausalLoss instead of hardcoding them at each call site --
  the root cause of a real shape-mismatch bug found in a demo notebook, and
  the same pattern that caused the predict_ute_* bug above.
* models.get_propensity_state_conditional_means(model): since the propensity
  head is no longer a single pypress.keras.layers.PredictiveStateMeans layer,
  reconstructs the same per-state constants (post-sigmoid) from the model's
  propensity_logit_state_<k> layers -- the replacement for the old
  model.layers[-2].state_conditional_means sanity check.
* metrics.uniform_entropy_gen: mean Shannon entropy of predictive state
  weights across a batch.

Other fixes
-----------
* BiasOnly.get_config() dropped `units` (silently reverting to 1 on reload)
  and stored `bias_regularizer` raw instead of via
  tf.keras.regularizers.serialize/deserialize.
* The three bootstrap_* functions seeded their resampling RNG inconsistently
  (RandomState(0) vs RandomState(n_samples)); all three now take an explicit
  random_state: int = 0.
* README.md's and docs/inference.md's code examples called a nonexistent
  inference.predict_ate(...) and called utils.split_y_pred(preds) with the
  wrong arity/unpacking -- anyone copy-pasting the README example hit an
  immediate error. Fixed and verified end-to-end.
* recommended_callbacks renamed to get_default_callbacks; `monitor` is now a
  required argument instead of defaulting to "val_loss".
* Bumped pypress dependency v0.0.6 -> v0.2.3; Uniform regularizer's `l1`
  penalty argument renamed to `l2` to track pypress's L1 -> L2 entropy penalty
  change.
* Added CHANGELOG.md, backfilling v0.0.1-v0.0.13 and documenting v0.1.0.

Version bumped to 0.1.0 in pyproject.toml (breaking propensity-head shape
change). Re-ran and fixed the tracked demo/example notebooks against the new
API.

Co-Authored-By: Claude Sonnet 5 <noreply@anthropic.com>
Claude-Session: https://claude.ai/code/session_01Kk1BtnCBYT2P4cHJHoMU4q
@gmgeorg
gmgeorg merged commit 44dfef8 into main Aug 22, 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