from __future__ import annotations import torch from torch import nn from torch.nn import functional as F class ContrastiveEncoder(nn.Module): def __init__(self) -> None: super().__init__() self.features = nn.Sequential( nn.Conv2d(1, 12, kernel_size=3, padding=1), nn.GELU(), nn.Conv2d(12, 24, kernel_size=3, padding=1), nn.GELU(), nn.MaxPool2d(2), nn.Flatten(), nn.Linear(24 * 4 * 4, 64), nn.GELU(), ) self.projector = nn.Sequential( nn.Linear(64, 32), nn.GELU(), nn.Linear(32, 16), ) def encode(self, pixels: torch.Tensor) -> torch.Tensor: return self.features(pixels) def forward(self, pixels: torch.Tensor) -> torch.Tensor: return self.projector(self.encode(pixels)) def parameter_count(model: nn.Module) -> int: return sum(parameter.numel() for parameter in model.parameters()) def nt_xent(first: torch.Tensor, second: torch.Tensor, temperature: float) -> torch.Tensor: batch_size = len(first) representations = F.normalize(torch.cat([first, second]), dim=1) similarities = representations @ representations.T / temperature diagonal = torch.eye(2 * batch_size, dtype=torch.bool) similarities = similarities.masked_fill(diagonal, -1e9) positives = torch.cat( [ torch.arange(batch_size, 2 * batch_size), torch.arange(0, batch_size), ] ) return F.cross_entropy(similarities, positives)