diff --git a/optax/losses/_classification.py b/optax/losses/_classification.py index 95e09f32a..64d882c54 100644 --- a/optax/losses/_classification.py +++ b/optax/losses/_classification.py @@ -815,9 +815,17 @@ def loop_body(prev, x): logalpha_phi_last = update_phi_score(logalpha_phi[-1], logalpha_emit[-1]) logalpha_phi = logalpha_phi.at[-1].set(logalpha_phi_last) + # Probabilities cannot exceed 1.0 (log-probabilities cannot exceed 0.0). Due + # to float32 rounding in log_softmax and logaddexp, values can round slightly + # above zero for confident predictions. Clamp to maintain mathematical bounds. + logalpha_phi = jnp.minimum(logalpha_phi, 0.0) + logalpha_emit = jnp.minimum(logalpha_emit, 0.0) + logalpha_phi_last = logalpha_phi[-1] + # extract per_seq_loss one_hot = jax.nn.one_hot(labellens, num_classes=maxlabellen + 1) # [B, N+1] per_seq_loss = -jnp.einsum('bn,bn->b', logalpha_phi_last, one_hot) # pylint:disable=invalid-unary-operand-type + per_seq_loss = jnp.maximum(per_seq_loss, 0.0) return per_seq_loss, logalpha_phi, logalpha_emit diff --git a/optax/losses/_classification_test.py b/optax/losses/_classification_test.py index da2b8173a..bdc7fda0e 100644 --- a/optax/losses/_classification_test.py +++ b/optax/losses/_classification_test.py @@ -963,6 +963,38 @@ def test_repeat_with_one_to_one_alignment(self): jnp.array(expected_loss), per_seq_loss[n], rtol=self._rtol ) + def test_confident_alignment_nonnegative(self): + # Tests that CTC loss and forward log-probabilities do not violate their + # theoretical bounds due to float32 rounding for confident alignments. + # See https://github.com/google-deepmind/optax/issues/1771 + logits = jnp.array([[[0.0, 17.0], [0.0, 17.0]]], dtype=jnp.float32) + loss, logalpha_phi, logalpha_emit = jax.jit( + _classification.ctc_loss_with_forward_probs, + static_argnames=('blank_id',), + )( + logits=logits, + logit_paddings=jnp.zeros((1, 2)), + labels=jnp.array([[1]], dtype=jnp.int32), + label_paddings=jnp.zeros((1, 1)), + blank_id=0, + ) + self.assertTrue(jnp.all(loss >= 0.0)) + self.assertTrue(jnp.all(logalpha_phi <= 0.0)) + self.assertTrue(jnp.all(logalpha_emit <= 0.0)) + + loss_direct = jax.jit( + _classification.ctc_loss, + static_argnames=('blank_id',), + )( + logits=logits, + logit_paddings=jnp.zeros((1, 2)), + labels=jnp.array([[1]], dtype=jnp.int32), + label_paddings=jnp.zeros((1, 1)), + blank_id=0, + ) + self.assertTrue(jnp.all(loss_direct >= 0.0)) + np.testing.assert_allclose(loss_direct, 0.0, atol=1e-6) + class SigmoidFocalLossTest(parameterized.TestCase):