| 74 | |
| 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 |
| 104 | ########### |
| 105 | def build_model(self): |
| 106 | raise NotImplementedError |
| 107 | |
| 108 | def load_ckpt(self, ckpt_base_dir, current_model_name=None, model_name='model', force=True, strict=True): |
| 109 | # This function is updated on 2021.12.13 |
| 110 | if current_model_name is None: |
| 111 | current_model_name = model_name |
| 112 | utils.load_ckpt(self.__getattr__(current_model_name), ckpt_base_dir, current_model_name, force, strict) |
| 113 | |
| 114 | def on_epoch_start(self): |
| 115 | self.training_losses_meter = {'total_loss': utils.AvgrageMeter()} |
| 116 | |
| 117 | def _training_step(self, sample, batch_idx, optimizer_idx): |
| 118 | """ |
| 119 | |
| 120 | :param sample: |
| 121 | :param batch_idx: |
| 122 | :return: total loss: torch.Tensor, loss_log: dict |
| 123 | """ |
| 124 | raise NotImplementedError |
| 125 | |
| 126 | def training_step(self, sample, batch_idx, optimizer_idx=-1): |
| 127 | loss_ret = self._training_step(sample, batch_idx, optimizer_idx) |
| 128 | self.opt_idx = optimizer_idx |
| 129 | if loss_ret is None: |
| 130 | return {'loss': None} |
| 131 | total_loss, log_outputs = loss_ret |
| 132 | log_outputs = utils.tensors_to_scalars(log_outputs) |
| 133 | for k, v in log_outputs.items(): |
nothing calls this directly
no outgoing calls
no test coverage detected