Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
8 changes: 8 additions & 0 deletions optax/losses/_classification.py
Original file line number Diff line number Diff line change
Expand Up @@ -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

Expand Down
32 changes: 32 additions & 0 deletions optax/losses/_classification_test.py
Original file line number Diff line number Diff line change
Expand Up @@ -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):

Expand Down
Loading