Skip to content

Stabilize NT-Xent loss and gradients at low temperatures - #1798

Open
Ayushdevo wants to merge 1 commit into
google-deepmind:mainfrom
Ayushdevo:fix/ntxent-logspace
Open

Ayushdevo wants to merge 1 commit into
google-deepmind:mainfrom
Ayushdevo:fix/ntxent-logspace

Conversation

@Ayushdevo

Copy link
Copy Markdown

At low positive temperatures, ntxent can return a finite loss with NaN gradients, or even negative infinity for batches containing multiple positive pairs. The current shared shift can underflow individual pair denominators to zero.

Reproduction on CPU (JAX/jaxlib 0.11.2):

import jax
import jax.numpy as jnp
import optax
x = jnp.array([[1., 0.], [0.5, 1.], [-1., 0.], [0., -1.]])
for labels in ([0, 0, 1, 1], [0, 0, 0, 1], [0, 0, 0, 0]):
    value, grad = jax.value_and_grad(optax.ntxent)(
        x, jnp.array(labels), 0.001)
    print(value, jnp.isfinite(grad).all())
# Before: 0.1732868 False; -inf False; -inf False

Compute the negative partition with logsumexp and each positive-pair loss with logaddexp. Guard rows without negatives before the reduction so their backward pass remains finite. This preserves the existing positive-pair-versus-all-negatives semantics and quadratic memory usage; batches without any positive pairs remain outside this change.

Add 12 regression cases comparing values and embedding gradients against an independent pair-by-pair reference, covering ordinary/low temperatures, two positive groups, a singleton group, all-same labels, and eager/JIT execution. Six cases fail on the original implementation.

Validation:

  • python -m pytest optax/losses -q: 190 passed, 6 subtests passed (10 float64-truncation warnings with x64 disabled).
  • Ruff on the two changed files, license-header check, and git diff --check pass.
  • Full repository suite not run.

The earlier #946 addressed zero embedding normalization; this change addresses denominator underflow in the contrastive loss itself.

Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Labels

None yet

Projects

None yet

Development

Successfully merging this pull request may close these issues.

1 participant