| 75 | |
| 76 | class BaseTask(nn.Module): |
| 77 | def __init__(self, *args, **kwargs): |
| 78 | # dataset configs |
| 79 | super(BaseTask, self).__init__(*args, **kwargs) |
| 80 | self.current_epoch = 0 |
| 81 | self.global_step = 0 |
| 82 | self.loaded_optimizer_states_dict = {} |
| 83 | self.trainer = None |
| 84 | self.logger = None |
| 85 | self.on_gpu = False |
| 86 | self.use_dp = False |
| 87 | self.use_ddp = False |
| 88 | self.example_input_array = None |
| 89 | |
| 90 | self.max_tokens = hparams['max_tokens'] |
| 91 | self.max_sentences = hparams['max_sentences'] |
| 92 | self.max_eval_tokens = hparams['max_eval_tokens'] |
| 93 | if self.max_eval_tokens == -1: |
| 94 | hparams['max_eval_tokens'] = self.max_eval_tokens = self.max_tokens |
| 95 | self.max_eval_sentences = hparams['max_eval_sentences'] |
| 96 | if self.max_eval_sentences == -1: |
| 97 | hparams['max_eval_sentences'] = self.max_eval_sentences = self.max_sentences |
| 98 | |
| 99 | self.model = None |
| 100 | self.training_losses_meter = None |
| 101 | |
| 102 | ########### |
| 103 | # Training, validation and testing |