Skip to content
Merged
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
1 change: 1 addition & 0 deletions .github/workflows/ci.yml
Original file line number Diff line number Diff line change
@@ -1,18 +1,18 @@
name: CI

Check warning on line 1 in .github/workflows/ci.yml

View workflow job for this annotation

GitHub Actions / zizmor-output

excessive-permissions

ci.yml:1: overly broad permissions: default permissions used due to no permissions: block
on: [push, pull_request]
jobs:
ruff:

Check warning on line 4 in .github/workflows/ci.yml

View workflow job for this annotation

GitHub Actions / zizmor-output

excessive-permissions

ci.yml:4: overly broad permissions: default permissions used due to no permissions: block
runs-on: ubuntu-latest
steps:
- uses: actions/checkout@v6

Check failure on line 7 in .github/workflows/ci.yml

View workflow job for this annotation

GitHub Actions / zizmor-output

zizmor/unpinned-uses

unpinned action reference: action is not pinned to a hash (required by blanket policy)

Check failure on line 7 in .github/workflows/ci.yml

View workflow job for this annotation

GitHub Actions / zizmor-output

unpinned-uses

ci.yml:7: unpinned action reference: action is not pinned to a hash (required by blanket policy)
- name: Lint
uses: astral-sh/ruff-action@v2

Check failure on line 9 in .github/workflows/ci.yml

View workflow job for this annotation

GitHub Actions / zizmor-output

zizmor/unpinned-uses

unpinned action reference: action is not pinned to a hash (required by blanket policy)

Check failure on line 9 in .github/workflows/ci.yml

View workflow job for this annotation

GitHub Actions / zizmor-output

unpinned-uses

ci.yml:9: unpinned action reference: action is not pinned to a hash (required by blanket policy)
test:

Check warning on line 10 in .github/workflows/ci.yml

View workflow job for this annotation

GitHub Actions / zizmor-output

excessive-permissions

ci.yml:10: overly broad permissions: default permissions used due to no permissions: block
runs-on: ubuntu-latest
steps:
- uses: actions/checkout@v6

Check failure on line 13 in .github/workflows/ci.yml

View workflow job for this annotation

GitHub Actions / zizmor-output

zizmor/unpinned-uses

unpinned action reference: action is not pinned to a hash (required by blanket policy)

Check failure on line 13 in .github/workflows/ci.yml

View workflow job for this annotation

GitHub Actions / zizmor-output

unpinned-uses

ci.yml:13: unpinned action reference: action is not pinned to a hash (required by blanket policy)
- name: Set up Python 3.12
uses: actions/setup-python@v6

Check failure on line 15 in .github/workflows/ci.yml

View workflow job for this annotation

GitHub Actions / zizmor-output

zizmor/unpinned-uses

unpinned action reference: action is not pinned to a hash (required by blanket policy)

Check failure on line 15 in .github/workflows/ci.yml

View workflow job for this annotation

GitHub Actions / zizmor-output

unpinned-uses

ci.yml:15: unpinned action reference: action is not pinned to a hash (required by blanket policy)
with:
python-version: 3.12
- name: Install dependencies
Expand All @@ -20,6 +20,7 @@
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/
Expand Down
10 changes: 4 additions & 6 deletions src/metrax/classification_metrics.py
Original file line number Diff line number Diff line change
Expand Up @@ -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


Expand Down
40 changes: 19 additions & 21 deletions src/metrax/nlp_metrics.py
Original file line number Diff line number Diff line change
Expand Up @@ -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:
Expand Down Expand Up @@ -216,20 +216,23 @@ 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_{<i})\right)
Perplexity = \exp\left(-\frac{1}{N}\sum_{i=1}^{N} \log
P(x_i|x_{<i})\right)

When sample weights :math:`w_i` are provided:

.. math::
Perplexity = \exp\left(-\frac{\sum_{i=1}^{N} w_i\log P(x_i|x_{<i})}{\sum_{i=1}^{N} w_i}\right)
Perplexity = \exp\left(-\frac{\sum_{i=1}^{N} w_i\log
P(x_i|x_{<i})}{\sum_{i=1}^{N} w_i}\right)

where:
- :math:`P(x_i|x_{<i})` is the predicted probability of token :math:`x_i`
given previous tokens
- :math:`w_i` are sample weights
- :math:`N` is the sequence length

Lower perplexity indicates better prediction - the model is less "perplexed" by the data.
Lower perplexity indicates better prediction - the model is less "perplexed"
by the data.
"""

aggregate_crossentropy: jax.Array
Expand All @@ -238,8 +241,9 @@ class Perplexity(clu_metrics.Metric):
@classmethod
def empty(cls) -> '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(
Expand All @@ -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.
Expand All @@ -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':
Expand Down Expand Up @@ -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:
Expand Down
Loading