(self, z1, z2)
| 64 | else: |
| 65 | return torch.mm(x, y) |
| 66 | def sim(self, z1, z2): |
| 67 | z1 = F.normalize(z1) |
| 68 | z2 = F.normalize(z2) |
| 69 | return torch.mm(z1, z2.t()) |
| 70 | |
| 71 | def batched_contrastive_loss(self, z1, z2, batch_size=4096): |
| 72 | device = z1.device |
no test coverage detected