Skip to content

Keep masked nan and inf out of nnx.metrics.Average - #5619

Open
vineethsaivs wants to merge 1 commit into
google:mainfrom
vineethsaivs:fix/nnx-average-mask-nonfinite
Open

vineethsaivs wants to merge 1 commit into
google:mainfrom
vineethsaivs:fix/nnx-average-mask-nonfinite

Conversation

@vineethsaivs

Copy link
Copy Markdown

nnx.metrics.Average.update drops masked entries with values * mask. A masked slot that holds nan or inf still poisons the total, since 0 * nan and 0 * inf are nan, so one padded position makes the average nan. That is common with per-token losses, where padding can sit over -inf logits. Accuracy and MultiMetric go through the same update.

The fix uses jnp.where(mask, values * mask, 0), so masked entries never reach the total. The count is unchanged.

Test: test_average_mask_skips_nonfinite puts nan and inf under a zero mask. On main it fails (total nan); with this change it passes, and tests/nnx/metrics_test.py passes 20/20 with JAX_NUMPY_RANK_PROMOTION=raise. ruff 0.1.3 and pyupgrade are clean.

This is separate from #5588 (count for broadcast masks). The two touch adjacent lines, so whichever lands second needs a small rebase.

Average.update zeroed masked entries with values * mask, but 0 * nan and
0 * inf are nan, so one non-finite value in a masked position made the
average nan. Drop masked entries with jnp.where instead.

This branch has not been deployed

No deployments
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