import torch
import torch.nn as nn
import torch.nn.functional as F

class ContrastiveHead(nn.Module):
    def __init__(self, in_dim=2048, proj_dim=128):
        super().__init__()
        self.mlp = nn.Sequential(
            nn.Linear(in_dim, in_dim),
            nn.BatchNorm1d(in_dim),
            nn.ReLU(inplace=True),
            nn.Linear(in_dim, proj_dim),
            nn.BatchNorm1d(proj_dim)
        )
    def forward(self, x):
        return F.normalize(self.mlp(x), dim=1)
