| 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 | """ |