MCPcopy Create free account
hub / github.com/MaureenZOU/TSAM / train

Method train

src/base/base_inference.py:93–135  ·  view source on GitHub ↗

Full training logic

(self)

Source from the content-addressed store, hash-verified

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):

Callers

nothing calls this directly

Calls 3

_train_epochMethod · 0.95
_save_checkpointMethod · 0.95
add_entryMethod · 0.80

Tested by

no test coverage detected