| 2 | from loguru import logger as loguru_logger |
| 3 | |
| 4 | class Logger: |
| 5 | def __init__(self, model, scheduler, cfg): |
| 6 | self.model = model |
| 7 | self.scheduler = scheduler |
| 8 | self.total_steps = 0 |
| 9 | self.running_loss = {} |
| 10 | self.writer = None |
| 11 | self.cfg = cfg |
| 12 | |
| 13 | def _print_training_status(self): |
| 14 | metrics_data = [self.running_loss[k]/self.cfg.sum_freq for k in sorted(self.running_loss.keys())] |
| 15 | training_str = "[{:6d}, {}] ".format(self.total_steps+1, self.scheduler.get_last_lr()) |
| 16 | metrics_str = ("{:10.4f}, "*len(metrics_data)).format(*metrics_data) |
| 17 | |
| 18 | # print the training status |
| 19 | loguru_logger.info(training_str + metrics_str) |
| 20 | |
| 21 | if self.writer is None: |
| 22 | if self.cfg.log_dir is None: |
| 23 | self.writer = SummaryWriter() |
| 24 | else: |
| 25 | self.writer = SummaryWriter(self.cfg.log_dir) |
| 26 | |
| 27 | for k in self.running_loss: |
| 28 | self.writer.add_scalar(k, self.running_loss[k]/self.cfg.sum_freq, self.total_steps) |
| 29 | self.running_loss[k] = 0.0 |
| 30 | |
| 31 | def push(self, metrics): |
| 32 | self.total_steps += 1 |
| 33 | |
| 34 | for key in metrics: |
| 35 | if key not in self.running_loss: |
| 36 | self.running_loss[key] = 0.0 |
| 37 | |
| 38 | self.running_loss[key] += metrics[key] |
| 39 | |
| 40 | if self.total_steps % self.cfg.sum_freq == self.cfg.sum_freq-1: |
| 41 | self._print_training_status() |
| 42 | self.running_loss = {} |
| 43 | |
| 44 | def write_dict(self, results): |
| 45 | if self.writer is None: |
| 46 | self.writer = SummaryWriter() |
| 47 | |
| 48 | for key in results: |
| 49 | self.writer.add_scalar(key, results[key], self.total_steps) |
| 50 | |
| 51 | def close(self): |
| 52 | self.writer.close() |
| 53 | |