(self, im_1, im_2)
| 136 | return F.cross_entropy(logits, labels) |
| 137 | |
| 138 | def forward(self, im_1, im_2): |
| 139 | feat_1 = self.encoder(im_1) # queries: NxC |
| 140 | proj_1 = self.projector(feat_1) |
| 141 | pred_1 = self.predictor(proj_1) |
| 142 | pred_1 = F.normalize(pred_1, dim=1) |
| 143 | |
| 144 | feat_2 = self.encoder(im_2) |
| 145 | proj_2 = self.projector(feat_2) |
| 146 | pred_2 = self.predictor(proj_2) |
| 147 | pred_2 = F.normalize(pred_2, dim=1) |
| 148 | |
| 149 | # compute key features |
| 150 | with torch.no_grad(): # no gradient to keys |
| 151 | self._momentum_update_key_encoder() # update the key encoder |
| 152 | |
| 153 | feat_1_ng = self.encoder_k(im_1) # keys: NxC |
| 154 | proj_1_ng = self.projector_k(feat_1_ng) |
| 155 | proj_1_ng = F.normalize(proj_1_ng, dim=1) |
| 156 | |
| 157 | feat_2_ng = self.encoder_k(im_2) |
| 158 | proj_2_ng = self.projector_k(feat_2_ng) |
| 159 | proj_2_ng = F.normalize(proj_2_ng, dim=1) |
| 160 | |
| 161 | # compute loss |
| 162 | loss = self.contrastive_loss(pred_1, proj_2_ng, self.queue2) \ |
| 163 | + self.contrastive_loss(pred_2, proj_1_ng, self.queue1) |
| 164 | |
| 165 | self._dequeue_and_enqueue(proj_1_ng, proj_2_ng) |
| 166 | |
| 167 | return loss |
| 168 | |
| 169 | |
| 170 | class MoBYMLP(nn.Module): |
nothing calls this directly
no test coverage detected