| 17 | return self.lm_head(logits) |
| 18 | |
| 19 | class GPT2Model(BaseModel): |
| 20 | def __init__(self, args, transformer=None, **kwargs): |
| 21 | super(GPT2Model, self).__init__(args, transformer=transformer, activation_func=gelu, **kwargs) |
| 22 | self.add_mixin("gpt2-final", GPT2FinalMixin(args.vocab_size, args.hidden_size)) |
| 23 | |
| 24 | @classmethod |
| 25 | def add_model_specific_args(cls, parser): |
| 26 | group = parser.add_argument_group('GPT2', 'GPT2 Configurations') |
| 27 | # group.add_argument('--num-types', type=int) |
| 28 | return parser |