Full training logic
(self)
| 91 | return device, list_ids |
| 92 | |
| 93 | def train(self): |
| 94 | """ |
| 95 | Full training logic |
| 96 | """ |
| 97 | epoch = 0 |
| 98 | result = self._train_epoch(epoch) |
| 99 | |
| 100 | # save logged informations into log dict |
| 101 | log = {'epoch': epoch} |
| 102 | for key, value in result.items(): |
| 103 | if key == 'metrics': |
| 104 | log.update({mtr.__name__: value[i] for i, mtr in enumerate(self.metrics)}) |
| 105 | elif key == 'val_metrics': |
| 106 | log.update({'val_' + mtr.__name__: value[i] for i, mtr in enumerate(self.metrics)}) |
| 107 | else: |
| 108 | log[key] = value |
| 109 | |
| 110 | # print logged informations to the screen |
| 111 | if self.train_logger is not None: |
| 112 | self.train_logger.add_entry(log) |
| 113 | if self.verbosity >= 1: |
| 114 | for key, value in log.items(): |
| 115 | self.logger.info(' {:15s}: {}'.format(str(key), value)) |
| 116 | |
| 117 | # evaluate model performance according to configured metric, save best checkpoint as model_best |
| 118 | best = False |
| 119 | monitor_value = None |
| 120 | if self.monitor_mode != 'off': |
| 121 | try: |
| 122 | if (self.monitor_mode == 'min' and log[self.monitor] < self.monitor_best) or\ |
| 123 | (self.monitor_mode == 'max' and log[self.monitor] > self.monitor_best): |
| 124 | self.monitor_best = log[self.monitor] |
| 125 | best = True |
| 126 | monitor_value = log[self.monitor] |
| 127 | |
| 128 | except KeyError: |
| 129 | if epoch == 1: |
| 130 | msg = "Warning: Can\'t recognize metric named '{}' ".format(self.monitor)\ |
| 131 | + "for performance monitoring. model_best checkpoint won\'t be updated." |
| 132 | self.logger.warning(msg) |
| 133 | |
| 134 | if epoch % self.save_freq == 0 or best: |
| 135 | self._save_checkpoint(epoch, save_best=best, monitor_value=monitor_value) |
| 136 | |
| 137 | |
| 138 | def _train_epoch(self, epoch): |
nothing calls this directly
no test coverage detected