From d0758d508e65b2a459e449de5517dc372193129e Mon Sep 17 00:00:00 2001 From: Tai An Date: Wed, 9 Sep 2026 00:17:47 -0700 Subject: [PATCH] [JAX] Fix undefined sr_rng_state in the encoder examples' --dry-run path test_model_parallel_encoder.py and test_multiprocessing_encoder.py build the dry-run rngs dict from `sr_rng_state`, which is never bound in train_and_evaluate; the local is `sr_rng`. Running either example with --dry-run raises NameError before the single train step. test_single_gpu_encoder.py and test_multigpu_encoder.py already use `sr_rng` in the same spot. Signed-off-by: Anai Guo --- examples/jax/encoder/test_model_parallel_encoder.py | 2 +- examples/jax/encoder/test_multiprocessing_encoder.py | 2 +- 2 files changed, 2 insertions(+), 2 deletions(-) diff --git a/examples/jax/encoder/test_model_parallel_encoder.py b/examples/jax/encoder/test_model_parallel_encoder.py index 4400485f26..27e2381fe7 100644 --- a/examples/jax/encoder/test_model_parallel_encoder.py +++ b/examples/jax/encoder/test_model_parallel_encoder.py @@ -382,7 +382,7 @@ def train_and_evaluate(args): if args.dry_run: labels = jnp.zeros(label_shape, dtype=jnp.bfloat16) - rngs = {DROPOUT_KEY: dropout_rng, SR_KEY: sr_rng_state} + rngs = {DROPOUT_KEY: dropout_rng, SR_KEY: sr_rng} jit_train_step(state, inputs, masks, labels, var_collect, rngs) print("PASSED") return None diff --git a/examples/jax/encoder/test_multiprocessing_encoder.py b/examples/jax/encoder/test_multiprocessing_encoder.py index 344e7d618b..c433b15874 100644 --- a/examples/jax/encoder/test_multiprocessing_encoder.py +++ b/examples/jax/encoder/test_multiprocessing_encoder.py @@ -479,7 +479,7 @@ def train_and_evaluate(args): if args.dry_run: labels = jnp.zeros(label_shape, dtype=jnp.bfloat16) - rngs = {DROPOUT_KEY: dropout_rng, SR_KEY: sr_rng_state} + rngs = {DROPOUT_KEY: dropout_rng, SR_KEY: sr_rng} jit_train_step(state, inputs, masks, labels, var_collect, rngs) print("PASSED") else: