MCPcopy Create free account
hub / github.com/drinkingcoder/FlowFormer-Official / Logger

Class Logger

core/utils/logger.py:4–52  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

2from loguru import logger as loguru_logger
3
4class 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

Callers 1

trainFunction · 0.90

Calls

no outgoing calls

Tested by

no test coverage detected