MCPcopy Create free account
hub / github.com/LayneH/GreenMIM / forward_decoder

Method forward_decoder

modeling/base_green_models.py:192–212  ·  view source on GitHub ↗
(self, x, ids_restore)

Source from the content-addressed store, hash-verified

190 return latent, mask, ids_restore
191
192 def forward_decoder(self, x, ids_restore):
193 # embed tokens
194 x = self.decoder_embed(x)
195
196 # append mask tokens to sequence
197 mask_tokens = self.mask_token.repeat(x.shape[0], ids_restore.shape[1] - x.shape[1], 1)
198 x_ = torch.cat([x, mask_tokens], dim=1)
199 x_ = torch.gather(x_, dim=1, index=ids_restore.unsqueeze(-1).repeat(1, 1, x.shape[2])) # unshuffle
200
201 # add pos embed
202 x = x_ + self.decoder_pos_embed
203
204 # apply Transformer blocks
205 for blk in self.decoder_blocks:
206 x = blk(x)
207 x = self.decoder_norm(x)
208
209 # predictor projection
210 x = self.decoder_pred(x)
211
212 return x
213
214 def forward_loss(self, imgs, pred, mask):
215 """

Callers 1

forwardMethod · 0.95

Calls

no outgoing calls

Tested by

no test coverage detected