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

class ContrastiveLoss(nn.Module):
    def __init__(self, temperature=0.7, neg_mode='all'):
        super().__init__(),
        self.temperature = temperature
        self.neg_mode = neg_mode
        self.cross_entropy = nn.CrossEntropyLoss()

    def forward(self, feat1, feat2):
        batch_size = feat1.size(0)
        sim_matrix = torch.einsum('id, jd->ij', feat1, feat2) / self.temperature
        labels = torch.arange(batch_size).to(feat1.device)
        if self.neg_mode == 'intra':
            mask = torch.eye(batch_size, dtype = torch.bool).to(feat1.device)
            sim_matrix = sim_matrix.masked_fill(mask, -1e9)
        
        loss = (self.cross_entropy(sim_matrix, labels) + self.cross_entropy(sim_matrix.T, labels))*0.5
        return loss

# class ProjectedContrastiveLoss(nn.Module):
#     def __init__(self, contrastive_head):
#         super().__init__()
#         self.projection = contrastive_head
#         self.criterion = ContrastiveLoss(temperature=temperature)

#     def forward(self, feat1, feat2):
#         if feat1.dim() == 4:
#             feat1 = torch.mean(feat1, dim=(2, 3))
#             feat2 = torch.mean(feat2, dim=(2, 3))
#         proj_feat1 = self.projection(feat1)
#         proj_feat2 = self.projection(feat2)
#         return self.criterion(proj_feat1, proj_feat2)

class ProjectedContrastiveLoss(nn.Module):
    def __init__(self, temperature=0.7, loss_weight=1.0):
        super().__init__()
        self.temperature = temperature
        self.loss_weight = loss_weight   
        self.cross_entropy = nn.CrossEntropyLoss()

    def forward(self, feat1, feat2):
        sim_matrix = torch.einsum('nc,mc->nm', feat1, feat2) / self.temperature
        labels = torch.arange(feat1.size(0)).to(feat1.device)
        loss = 0.5 * (self.cross_entropy(sim_matrix, labels) + self.cross_entropy(sim_matrix.T, labels))
        return loss * self.loss_weight