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: