Skip to content

fix: use y.shape[1] in get_n_cols for Keras 3 compatibility (v0.1.4) - #24

Merged
gmgeorg merged 1 commit into
mainfrom
fix/keras3-get-n-cols
Aug 23, 2026
Merged

gmgeorg merged 1 commit into
mainfrom
fix/keras3-get-n-cols

Conversation

@gmgeorg

@gmgeorg gmgeorg commented Aug 23, 2026

Copy link
Copy Markdown
Owner

Summary

  • get_n_cols (pypsps/utils.py) used the old TF1-era y.get_shape().as_list()[1] for non-np.ndarray inputs. Under Keras 3, the tensor flowing through the nested CausalLoss -> OutcomeLoss call chain during graph-mode loss tracing is a KerasTensor, which has no get_shape() method at all -- raising AttributeError and crashing training with the default (posterior-weighted) loss.
  • Fixed by using y.shape[1] unconditionally, which works uniformly for np.ndarray, eager tf.Tensor, graph-traced tensors, and KerasTensor, so the isinstance branch is no longer needed.
  • Bumped pyproject.toml to 0.1.4 and added a CHANGELOG.md entry.

Test plan

  • Added test_get_n_cols_np_array, test_get_n_cols_eager_tensor, test_get_n_cols_graph_mode_tensor, test_get_n_cols_keras_symbolic_tensor covering each tensor type get_n_cols needs to handle.
  • Added test_toy_model_fits_in_eager_and_graph_mode, which fits build_toy_model with run_eagerly=True and run_eagerly=False.
  • Verified against the pre-fix code: test_get_n_cols_keras_symbolic_tensor fails with AttributeError: 'KerasTensor' object has no attribute 'get_shape', confirming this reproduces the reported crash.
  • Full suite passes: pytest pypsps/tests/ -> 89 passed.

🤖 Generated with Claude Code

y.get_shape().as_list()[1] is old TF1-era API. Under Keras 3, the
tensor flowing through the nested CausalLoss -> OutcomeLoss call
chain during graph-mode loss tracing is a KerasTensor, which has no
get_shape() method at all, raising AttributeError and crashing
training with the default (posterior-weighted) loss.

Fixed by using y.shape[1] unconditionally, which works for
np.ndarray, eager tf.Tensor, graph-traced tensors, and KerasTensor
alike -- so the isinstance branch is no longer needed either.

Added regression tests covering all four tensor types plus an
integration test that fits the toy model with run_eagerly True and
False. Verified against the pre-fix code that the KerasTensor case
fails with the exact reported AttributeError.

Co-Authored-By: Claude Sonnet 5 <noreply@anthropic.com>
@gmgeorg
gmgeorg merged commit 62eaf3e into main Aug 23, 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