From 2ab9c954a880d79f3e853514e80e26974f431822 Mon Sep 17 00:00:00 2001 From: Metrax Authors Date: Tue, 15 Sep 2026 15:56:49 -0700 Subject: [PATCH] Fix Perplexity metric accumulation across tokens and sample weights in Metrax. Previously, `Perplexity.from_model_output` normalized cross-entropy per batch and weighted each batch by sequence count (`labels.shape[0]`) rather than total token count/weight. A recent KerasHub update (cl/974631537) changed the Perplexity metric calculation logic to normalize by total token count, which led to failing tests in Metrax. I have also fixed couple of legacy lint errors, because they were blocking presubmit. Updating the ci.yaml to use the keras-hub-nightly for testing in CI environment. PiperOrigin-RevId: 982091649 --- .github/workflows/ci.yml | 1 + src/metrax/classification_metrics.py | 10 +++---- src/metrax/nlp_metrics.py | 40 +++++++++++++--------------- 3 files changed, 24 insertions(+), 27 deletions(-) diff --git a/.github/workflows/ci.yml b/.github/workflows/ci.yml index b6dda03..f82bd92 100644 --- a/.github/workflows/ci.yml +++ b/.github/workflows/ci.yml @@ -20,6 +20,7 @@ jobs: python -m pip install --upgrade pip pip install torch --index-url https://download.pytorch.org/whl/cpu pip install ".[dev]" + pip install --upgrade keras-hub-nightly - name: Run Unit Tests run: | pytest ./src/ diff --git a/src/metrax/classification_metrics.py b/src/metrax/classification_metrics.py index 8ec9e1d..dbd12ff 100644 --- a/src/metrax/classification_metrics.py +++ b/src/metrax/classification_metrics.py @@ -67,12 +67,10 @@ def _squeeze_mismatching_trailing_ones( predictions: jax.Array, labels: jax.Array ) -> tuple[jax.Array, jax.Array]: """Squeezes mismatching trailing ones from predictions and labels.""" - if predictions.ndim < labels.ndim: - if labels.shape[-1] == 1: - labels = jnp.squeeze(labels, axis=-1) - elif labels.ndim < predictions.ndim: - if predictions.shape[-1] == 1: - predictions = jnp.squeeze(predictions, axis=-1) + if predictions.ndim < labels.ndim and labels.shape[-1] == 1: + labels = jnp.squeeze(labels, axis=-1) + elif labels.ndim < predictions.ndim and predictions.shape[-1] == 1: + predictions = jnp.squeeze(predictions, axis=-1) return predictions, labels diff --git a/src/metrax/nlp_metrics.py b/src/metrax/nlp_metrics.py index 7a99b70..7bca3a7 100644 --- a/src/metrax/nlp_metrics.py +++ b/src/metrax/nlp_metrics.py @@ -56,7 +56,7 @@ def _get_ngrams(segment: list[str], max_order: int): def _lcs_length(str1: list[str], str2: list[str]) -> int: """Computes the length of the Longest Common Subsequence (LCS).""" - lengths = [[0 for j in range(len(str2) + 1)] for i in range(len(str1) + 1)] + lengths = [[0 for _ in range(len(str2) + 1)] for _ in range(len(str1) + 1)] for i, x in enumerate(str1): for j, y in enumerate(str2): if x == y: @@ -216,12 +216,14 @@ class Perplexity(clu_metrics.Metric): Given a sequence of :math:`N` tokens, perplexity is calculated as: .. math:: - Perplexity = \exp\left(-\frac{1}{N}\sum_{i=1}^{N} \log P(x_i|x_{ 'Perplexity': return cls( - aggregate_crossentropy=jnp.array(0, jnp.float32), - num_samples=jnp.array(0, jnp.float32)) + aggregate_crossentropy=jnp.array(0, jnp.float32), + num_samples=jnp.array(0, jnp.float32), + ) @classmethod def from_model_output( @@ -252,12 +256,11 @@ def from_model_output( """Updates the metric. Args: - predictions: A floating point tensor representing the prediction - generated from the model. The shape should be (batch_size, seq_len, - vocab_size). + predictions: A floating point tensor representing the prediction generated + from the model. The shape should be (batch_size, seq_len, vocab_size). labels: True value. The shape should be (batch_size, seq_len). - sample_weights: An optional tensor representing the - weight of each token. The shape should be (batch_size, seq_len). + sample_weights: An optional tensor representing the weight of each token. + The shape should be (batch_size, seq_len). from_logits: Whether the predictions are logits. If True, the predictions are converted to probabilities using a softmax. If False, all values outside of [0, 1] are clipped to 0 or 1. @@ -282,20 +285,15 @@ def from_model_output( labels_one_hot = jax.nn.one_hot(labels, predictions.shape[-1], axis=-1) crossentropy = -jnp.sum(labels_one_hot * log_prob, axis=-1) - # Sum across sequence length dimension first. if sample_weights is not None: + num_tokens = jnp.sum(sample_weights) crossentropy = crossentropy * sample_weights - # Normalize by the sum of weights for each sequence. - crossentropy = base.divide_no_nan( - jnp.sum(crossentropy), jnp.sum(sample_weights) - ) else: - crossentropy = jnp.mean(crossentropy) + num_tokens = jnp.array(labels.size) - batch_size = jnp.array(labels.shape[0]) return cls( - aggregate_crossentropy=(batch_size * crossentropy), - num_samples=batch_size, + aggregate_crossentropy=(jnp.sum(crossentropy)), + num_samples=num_tokens, ) def merge(self, other: 'Perplexity') -> 'Perplexity': @@ -669,7 +667,7 @@ def from_model_output( ) @staticmethod - def _levenshtein_distance(prediction: list, reference: list) -> int: + def _levenshtein_distance(prediction: list[str], reference: list[str]) -> int: """Computes the Levenshtein (edit) distance between two token sequences. Args: