Keras 3 determinism behavior across JAX and PyTorch backends
0 reputation · 19 Apr 2026, 14:07 UTC
Reproducible Execution in Multi-Backend Environments
Keras 3 provides keras.config.enable_determinism() to ensure consistent model behavior across different runs. While this configuration manages high-level settings, the underlying execution depends on the active backend (TensorFlow, JAX, or PyTorch).
There is uncertainty regarding how this global Keras flag interacts with backend-specific random number generators. For instance, JAX requires explicit PRNG keys, and PyTorch relies on its own manual seeding mechanisms to guarantee deterministic tensor operations.
When utilizing a repeatable development environment, it is unclear if enable_determinism(True) is sufficient to synchronize these low-level backend seeds or if manual integration with jax.random.key() and torch.manual_seed() remains mandatory for full reproducibility.
- Does
keras.config.enable_determinism()automatically trigger the corresponding seed settings in JAX and PyTorch? - What is the recommended verification method to ensure weight updates are identical across different backends under this setting?