Conversation
… (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
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.
Squashes the fix/1 branch's 8 commits into one for review. Summary of what changed and why:
Core correctness fix
ais 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.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 rebindingloss._alphabetween 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.alphais now a tf.Variable, updated via.assign(...). Verified directly (tf.function trace before/after) and end-to-end via model.fit().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
Other fixes
units(silently reverting to 1 on reload) and storedbias_regularizerraw instead of via tf.keras.regularizers.serialize/deserialize.monitoris now a required argument instead of defaulting to "val_loss".l1penalty argument renamed tol2to track pypress's L1 -> L2 entropy penalty change.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