| 122 | return logits[:, 1:] |
| 123 | |
| 124 | class MAEDecoder(BaseModel): |
| 125 | def __init__(self, args, transformer=None, layernorm_epsilon=1e-6): |
| 126 | super().__init__(args, transformer=transformer, layernorm_epsilon=layernorm_epsilon) |
| 127 | self.add_mixin('mask_forward', MaskMixin(args)) |
| 128 | @classmethod |
| 129 | def add_model_specific_args(cls, parser): |
| 130 | return super().add_model_specific_args(parser) |
| 131 | |
| 132 | from sat.model import EncoderDecoderModel |
| 133 | import argparse |