Source code for core.nn.metrics.rankme

import logging

import torch
import torch.nn.functional as F

logger = logging.getLogger(__name__)


[docs] def rankme(x: torch.Tensor, normalize: bool = True) -> float: """Effective rank of an ``(entities, dim)`` embedding: the entropy of its spectrum.""" assert x.ndim == 2 # (neurons, dim) invalid_emb = (~x.isfinite()).sum(dim=1) > 0 x = x[~invalid_emb, :] if normalize: x = F.normalize(x, dim=-1) try: cov = torch.cov(x.T).float() eigvals = torch.linalg.eigvalsh(cov) eigvals = eigvals / eigvals.sum() eigvals = eigvals[eigvals > 0] rankme = torch.exp(-(eigvals * torch.log(eigvals)).sum()) except Exception: logger.error("Error computing rankme", exc_info=True) rankme = torch.nan return rankme.item()