| 407 | return x, gt_indices, token_drop_mask, token_all_mask |
| 408 | |
| 409 | def forward_decoder(self, x, token_drop_mask, token_all_mask): |
| 410 | # embed tokens |
| 411 | x = self.decoder_embed(x) |
| 412 | |
| 413 | # append mask tokens to sequence |
| 414 | if self.pad_with_cls_token: |
| 415 | mask_tokens = x[:, 0:1].repeat(1, token_all_mask.shape[1], 1) |
| 416 | else: |
| 417 | mask_tokens = self.mask_token.repeat(token_all_mask.shape[0], token_all_mask.shape[1], 1) |
| 418 | |
| 419 | # put undropped tokens into original sequence |
| 420 | x_after_pad = mask_tokens.clone() |
| 421 | x_after_pad[(1 - token_drop_mask).nonzero(as_tuple=True)] = x.reshape(x.shape[0] * x.shape[1], x.shape[2]) |
| 422 | # set undropped but masked positions with mask |
| 423 | x_after_pad = torch.where(token_all_mask.unsqueeze(-1).bool(), mask_tokens, x_after_pad) |
| 424 | |
| 425 | # add pos embed |
| 426 | x = x_after_pad + self.decoder_pos_embed_learned |
| 427 | |
| 428 | # apply Transformer blocks |
| 429 | for blk in self.decoder_blocks: |
| 430 | x = blk(x) |
| 431 | |
| 432 | x = self.decoder_norm(x) |
| 433 | |
| 434 | word_embeddings = self.token_emb.word_embeddings.weight.data.detach() |
| 435 | x = self.mlm_layer(x, word_embeddings) |
| 436 | # print("Logits shape:", x.shape) |
| 437 | |
| 438 | return x |
| 439 | |
| 440 | def normalize(self, x): |
| 441 | return (x - torch.mean(x, dim=1, keepdim=True)) / torch.std(x, dim=1, keepdim=True) |