From 433bb1fb394cee3cecdb89327854664770ce6714 Mon Sep 17 00:00:00 2001 From: wahid18-maqs Date: Fri, 31 Jul 2026 11:31:10 +0000 Subject: [PATCH 1/2] fix: reset notfinite_count after apply_if_finite gives up --- optax/transforms/_conditionality.py | 7 +++++-- optax/transforms/_conditionality_test.py | 2 ++ 2 files changed, 7 insertions(+), 2 deletions(-) diff --git a/optax/transforms/_conditionality.py b/optax/transforms/_conditionality.py index a06b3a9a4..71cd2fae1 100644 --- a/optax/transforms/_conditionality.py +++ b/optax/transforms/_conditionality.py @@ -238,10 +238,13 @@ def update(updates, state, params=None, **extra_args): isfinite = jnp.all( jnp.array([jnp.all(jnp.isfinite(p)) for p in flat_updates]) ) + give_up = jnp.logical_and( + jnp.logical_not(isfinite), state.notfinite_count >= max_consecutive_errors + ) notfinite_count = jnp.where( isfinite, jnp.zeros([], jnp.int32), - numerics.safe_increment(state.notfinite_count), + jnp.where(give_up, jnp.zeros([], jnp.int32), numerics.safe_increment(state.notfinite_count)), ) def do_update(_): @@ -251,7 +254,7 @@ def reject_update(_): return optax.tree.zeros_like(updates), inner_state updates, new_inner_state = lax.cond( - jnp.logical_or(isfinite, notfinite_count > max_consecutive_errors), + jnp.logical_or(isfinite, give_up), do_update, reject_update, operand=None, diff --git a/optax/transforms/_conditionality_test.py b/optax/transforms/_conditionality_test.py index 31a642704..4d75a8afd 100644 --- a/optax/transforms/_conditionality_test.py +++ b/optax/transforms/_conditionality_test.py @@ -105,6 +105,7 @@ def fn(p, x): updates, state = opt.update(grads, state, params) params = update.apply_updates(params, updates) self.assertTrue(bool(jnp.isnan(jax.tree.flatten(params)[0][0]))) + self.assertEqual(0, int(getattr(state, 'notfinite_count'))) self.assertEqual(5, int(getattr(state, 'total_notfinite'))) def test_apply_if_finite_pmap(self): @@ -150,6 +151,7 @@ def fn_update(params, opt_state, x): self.assertEqual(step + 1, opt_state.notfinite_count.item()) # Next param update with NaN is accepted since we reached maximum _, opt_state = fn_update(params, opt_state, two) + self.assertEqual(0, opt_state.notfinite_count.item()) self.assertEqual(5, opt_state.total_notfinite.item()) From 2454bf96b2bc5cfdde309cfc6dd3cbcee1b3b607 Mon Sep 17 00:00:00 2001 From: wahid18-maqs Date: Fri, 31 Jul 2026 11:34:27 +0000 Subject: [PATCH 2/2] Fixed the line-length violation --- optax/transforms/_conditionality.py | 9 +++++++-- 1 file changed, 7 insertions(+), 2 deletions(-) diff --git a/optax/transforms/_conditionality.py b/optax/transforms/_conditionality.py index 71cd2fae1..edff46b7f 100644 --- a/optax/transforms/_conditionality.py +++ b/optax/transforms/_conditionality.py @@ -239,12 +239,17 @@ def update(updates, state, params=None, **extra_args): jnp.array([jnp.all(jnp.isfinite(p)) for p in flat_updates]) ) give_up = jnp.logical_and( - jnp.logical_not(isfinite), state.notfinite_count >= max_consecutive_errors + jnp.logical_not(isfinite), + state.notfinite_count >= max_consecutive_errors, ) notfinite_count = jnp.where( isfinite, jnp.zeros([], jnp.int32), - jnp.where(give_up, jnp.zeros([], jnp.int32), numerics.safe_increment(state.notfinite_count)), + jnp.where( + give_up, + jnp.zeros([], jnp.int32), + numerics.safe_increment(state.notfinite_count), + ), ) def do_update(_):