Conversation
This file contains hidden or bidirectional Unicode text that may be interpreted or compiled differently than what appears below. To review, open the file in an editor that reveals hidden Unicode characters.
Learn more about bidirectional Unicode characters
Sign up for free
to join this conversation on GitHub.
Already have an account?
Sign in to comment
Add this suggestion to a batch that can be applied as a single commit.This suggestion is invalid because no changes were made to the code.Suggestions cannot be applied while the pull request is closed.Suggestions cannot be applied while viewing a subset of changes.Only one suggestion per line can be applied in a batch.Add this suggestion to a batch that can be applied as a single commit.Applying suggestions on deleted lines is not supported.You must change the existing code in this line in order to create a valid suggestion.Outdated suggestions cannot be applied.This suggestion has been applied or marked resolved.Suggestions cannot be applied from pending reviews.Suggestions cannot be applied on multi-line comments.Suggestions cannot be applied while the pull request is queued to merge.Suggestion cannot be applied right now. Please check back later.
At low positive temperatures,
ntxentcan 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):
Compute the negative partition with
logsumexpand each positive-pair loss withlogaddexp. 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).git diff --checkpass.The earlier #946 addressed zero embedding normalization; this change addresses denominator underflow in the contrastive loss itself.