| 160 | return super().add_model_specific_args(parser) |
| 161 | |
| 162 | class CaiTDecoder(BaseModel): |
| 163 | def __init__(self, args, transformer=None, layernorm_epsilon=1e-6): |
| 164 | super().__init__(args, is_decoder=True, transformer=transformer, layernorm_epsilon=layernorm_epsilon) |
| 165 | self.add_mixin('cls', ClsMixin(args.hidden_size, args.num_classes)) |
| 166 | self.add_mixin('dec_forward', DecForward(args.hidden_size, args.num_layers, init_values=args.init_scale)) |
| 167 | @classmethod |
| 168 | def add_model_specific_args(cls, parser): |
| 169 | return super().add_model_specific_args(parser) |
| 170 | |
| 171 | from sat.model import EncoderDecoderModel |
| 172 | import argparse |