| 133 | import argparse |
| 134 | |
| 135 | class MAE(EncoderDecoderModel): |
| 136 | def __init__(self, args, transformer=None, layernorm_epsilon=1e-6): |
| 137 | encoder = MAEEncoder(args, transformer=transformer, layernorm_epsilon=layernorm_epsilon) |
| 138 | dec_args = argparse.Namespace(**vars(args)) |
| 139 | # dec_args.enc_hidden_size = dec_args.hidden_size # used for cross attn |
| 140 | override_attrs = ['num_layers', 'hidden_size', 'num_attention_heads', |
| 141 | 'max_sequence_length', 'inner_hidden_size', 'hidden_size_per_attention_head'] |
| 142 | for name in override_attrs: |
| 143 | dec_attr = getattr(dec_args, 'dec_' + name, None) |
| 144 | if dec_attr is not None: # else use encoder-config |
| 145 | setattr(dec_args, name, dec_attr) |
| 146 | setattr(dec_args, 'enc_hidden_size', args.hidden_size) |
| 147 | decoder = MAEDecoder(dec_args, transformer=transformer, layernorm_epsilon=layernorm_epsilon) |
| 148 | super().__init__(args, encoder=encoder, decoder=decoder, tie_word_embeddings=False) |
| 149 | |
| 150 | def encode(self, input_ids, position_ids, attention_mask=None, **kw_args): |
| 151 | return self.encoder(input_ids, position_ids, attention_mask, **kw_args) |
| 152 | |
| 153 | def decode(self, input_ids, position_ids, attention_mask, encoder_outputs, ids_restore, **kw_args): |
| 154 | return self.decoder(input_ids, position_ids, attention_mask, encoder_outputs=encoder_outputs, ids_restore=ids_restore, **kw_args) |
| 155 | |
| 156 | def forward(self, input_ids, enc_position_ids, dec_position_ids, *, enc_attention_mask=None, dec_attention_mask=None, **kw_args): |
| 157 | if enc_attention_mask is None: |
| 158 | enc_attention_mask = torch.ones(1, 1, dtype=self.encoder.transformer.word_embeddings.weight.dtype, device=input_ids.device) |
| 159 | encoder_outputs, *encoder_mems = self.encode(input_ids, enc_position_ids, enc_attention_mask, **kw_args) |
| 160 | decoder_outputs, *decoder_mems = self.decode(input_ids, dec_position_ids, dec_attention_mask, encoder_outputs=encoder_outputs, ids_restore=encoder_mems[0]["ids_restore"], **kw_args) |
| 161 | return encoder_outputs, decoder_outputs, encoder_mems, decoder_mems |
| 162 | |
| 163 | def unpatchify(self, x): |
| 164 | """ |
| 165 | x: (N, L, patch_size**2 *3) |
| 166 | imgs: (N, 3, H, W) |
| 167 | """ |
| 168 | p = self.encoder.property.patch_size |
| 169 | h = w = int(x.shape[1]**.5) |
| 170 | assert h * w == x.shape[1] |
| 171 | |
| 172 | x = x.reshape(shape=(x.shape[0], h, w, p, p, 3)) |
| 173 | x = torch.einsum('nhwpqc->nchpwq', x) |
| 174 | imgs = x.reshape(shape=(x.shape[0], 3, h * p, h * p)) |
| 175 | return imgs |