| 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) |
|
|