MCPcopy Create free account
hub / github.com/MoonInTheRiver/DiffSinger / BaseTask

Class BaseTask

tasks/base_task.py:76–359  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

74
75
76class 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():

Callers

nothing calls this directly

Calls

no outgoing calls

Tested by

no test coverage detected