| 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. |