MCPcopy Create free account
hub / github.com/SwinTransformer/Transformer-SSL / forward

Method forward

models/moby.py:138–167  ·  view source on GitHub ↗
(self, im_1, im_2)

Source from the content-addressed store, hash-verified

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
170class MoBYMLP(nn.Module):

Callers

nothing calls this directly

Calls 3

contrastive_lossMethod · 0.95
_dequeue_and_enqueueMethod · 0.95

Tested by

no test coverage detected