| 172 | import argparse |
| 173 | |
| 174 | class CaiT(EncoderDecoderModel): |
| 175 | def __init__(self, args, transformer=None, layernorm_epsilon=1e-6): |
| 176 | encoder = CaiTEncoder(args, transformer=transformer, layernorm_epsilon=layernorm_epsilon) |
| 177 | dec_args = argparse.Namespace(**vars(args)) |
| 178 | # dec_args.enc_hidden_size = dec_args.hidden_size # used for cross attn |
| 179 | override_attrs = ['num_layers', 'hidden_size', 'num_attention_heads', 'layernorm_order', |
| 180 | 'max_sequence_length', 'inner_hidden_size', 'hidden_size_per_attention_head'] |
| 181 | for name in override_attrs: |
| 182 | dec_attr = getattr(dec_args, 'dec_' + name, None) |
| 183 | if dec_attr is not None: # else use encoder-config |
| 184 | setattr(dec_args, name, dec_attr) |
| 185 | decoder = CaiTDecoder(dec_args, transformer=transformer, layernorm_epsilon=layernorm_epsilon) |
| 186 | super().__init__(args, encoder=encoder, decoder=decoder) |
| 187 | |
| 188 | def forward(self, input_ids, enc_position_ids, dec_position_ids, *, enc_attention_mask=None, dec_attention_mask=None, cross_attention_mask=None, **kw_args): |
| 189 | # Please use self.decoder for auto-regressive generation. |
| 190 | if enc_attention_mask is None: |
| 191 | enc_attention_mask = torch.ones(1, 1, dtype=self.encoder.transformer.word_embeddings.weight.dtype, device=input_ids.device) |
| 192 | if cross_attention_mask is None: |
| 193 | cross_attention_mask = enc_attention_mask |
| 194 | encoder_outputs = self.encode(input_ids, enc_position_ids, enc_attention_mask, **kw_args) |
| 195 | decoder_outputs, *mems = self.decode(input_ids, dec_position_ids, dec_attention_mask, encoder_outputs=encoder_outputs, cross_attention_mask=cross_attention_mask, **kw_args) |
| 196 | return (encoder_outputs, decoder_outputs, *mems) |