Skip to content

Natural Language Processing

tensoris.backend.metrics.nlp

Natural Language Processing Evaluation Metrics.

Classes

PerplexityMetric

Perplexity metric for language model evaluation.

Formula
\[\text{PPL} = \exp\left( -\frac{1}{N} \sum_{i=1}^N \log p(x_i \mid x_{<i}) \right) = \exp\left( \mathcal{L}_{\text{CE}} \right)\]

where \(\mathcal{L}_{\text{CE}}\) is the average cross-entropy loss per token.

References

Jelinek, F., Mercer, R. L., Bahl, L. R., & Baker, J. K. (1977). Perplexity—a measure of the difficulty of speech recognition tasks. Journal of the Acoustical Society of America, 62(S1), S63-S63.

Source code in src/tensoris/backend/metrics/nlp.py
class PerplexityMetric:
    """Perplexity metric for language model evaluation.

    Formula:
        $$\\text{PPL} = \\exp\\left( -\\frac{1}{N} \\sum_{i=1}^N \\log p(x_i \\mid x_{<i}) \\right) = \\exp\\left( \\mathcal{L}_{\\text{CE}} \\right)$$

        where $\\mathcal{L}_{\\text{CE}}$ is the average cross-entropy loss per token.

    References:
        Jelinek, F., Mercer, R. L., Bahl, L. R., & Baker, J. K. (1977).
        Perplexity—a measure of the difficulty of speech recognition tasks.
        Journal of the Acoustical Society of America, 62(S1), S63-S63.
    """

    def __init__(self) -> None:
        """Initialize perplexity accumulator."""
        self.reset()

    def reset(self) -> None:
        """Reset accumulated loss counts."""
        self.total_loss = 0.0
        self.total_count = 0

    def update(
        self, cross_entropy_loss: torch.Tensor | float, token_count: int
    ) -> None:
        """Accumulate token cross-entropy loss.

        Args:
            cross_entropy_loss: Average batch cross-entropy loss.
            token_count: Number of active non-padding tokens in batch.
        """
        loss_val = (
            cross_entropy_loss.item()
            if hasattr(cross_entropy_loss, "item")
            else float(cross_entropy_loss)
        )
        self.total_loss += loss_val * token_count
        self.total_count += token_count

    def compute(self) -> float:
        """Compute exponentiated mean cross-entropy perplexity.

        Returns:
            Scalar perplexity value exp(mean_loss).
        """
        if self.total_count == 0:
            return 0.0
        mean_loss = self.total_loss / self.total_count
        return torch.exp(torch.tensor(mean_loss)).item()
Methods:
__init__
__init__()

Initialize perplexity accumulator.

Source code in src/tensoris/backend/metrics/nlp.py
def __init__(self) -> None:
    """Initialize perplexity accumulator."""
    self.reset()
compute
compute()

Compute exponentiated mean cross-entropy perplexity.

Returns:

Type Description
float

Scalar perplexity value exp(mean_loss).

Source code in src/tensoris/backend/metrics/nlp.py
def compute(self) -> float:
    """Compute exponentiated mean cross-entropy perplexity.

    Returns:
        Scalar perplexity value exp(mean_loss).
    """
    if self.total_count == 0:
        return 0.0
    mean_loss = self.total_loss / self.total_count
    return torch.exp(torch.tensor(mean_loss)).item()
reset
reset()

Reset accumulated loss counts.

Source code in src/tensoris/backend/metrics/nlp.py
def reset(self) -> None:
    """Reset accumulated loss counts."""
    self.total_loss = 0.0
    self.total_count = 0
update
update(cross_entropy_loss, token_count)

Accumulate token cross-entropy loss.

Parameters:

Name Type Description Default
cross_entropy_loss Tensor | float

Average batch cross-entropy loss.

required
token_count int

Number of active non-padding tokens in batch.

required
Source code in src/tensoris/backend/metrics/nlp.py
def update(
    self, cross_entropy_loss: torch.Tensor | float, token_count: int
) -> None:
    """Accumulate token cross-entropy loss.

    Args:
        cross_entropy_loss: Average batch cross-entropy loss.
        token_count: Number of active non-padding tokens in batch.
    """
    loss_val = (
        cross_entropy_loss.item()
        if hasattr(cross_entropy_loss, "item")
        else float(cross_entropy_loss)
    )
    self.total_loss += loss_val * token_count
    self.total_count += token_count