def nt_xent_loss(x, temperature): assert len(x.size()) == 2 # Cosine similarity xcs = F.cosine_similarity(x[None,:,:], x[:,None,:], dim=-1) xcs[torch.eye(x.size(0)).bool()] = float("-inf") # Ground truth labels target = torch.arange(8) target[0::2] += 1 target[1::2] -= 1 # Standard cross-entropy loss return F.cross_entropy(xcs / temperature, target, reduction="mean")