MCPcopy Create free account
hub / github.com/LTH14/rcg / forward_decoder

Method forward_decoder

pixel_generator/mage/models_mage.py:409–438  ·  view source on GitHub ↗
(self, x, token_drop_mask, token_all_mask)

Source from the content-addressed store, hash-verified

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)

Callers 2

forwardMethod · 0.95
gen_imageMethod · 0.95

Calls

no outgoing calls

Tested by

no test coverage detected