(self, args, transformer=None, layernorm_epsilon=1e-6)
| 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) |
no test coverage detected