Coverage for hopwise/utils/wandblogger.py: 36%
36 statements
« prev ^ index » next coverage.py v7.16.2, created at 2026-09-30 13:25 +0000
« prev ^ index » next coverage.py v7.16.2, created at 2026-09-30 13:25 +0000
1# @Time : 2022/8/2
2# @Author : Ayush Thakur
3# @Email : ayusht@wandb.com
5r"""hopwise.utils.wandblogger
6################################
7"""
10class WandbLogger:
11 """WandbLogger to log metrics to Weights and Biases."""
13 def __init__(self, config):
14 """Args:
15 config (dict): A dictionary of parameters used by hopwise.
16 """
17 self.config = config
18 self.log_wandb = config.log_wandb
19 self.setup()
21 def setup(self):
22 if self.log_wandb:
23 try:
24 import wandb
26 self._wandb = wandb
27 except ImportError:
28 raise ImportError(
29 "To use the Weights and Biases Logger please install wandb.Run `pip install wandb` to install it."
30 )
32 # Initialize a W&B run
33 if self._wandb.run is None:
34 self._wandb.init(project=self.config.wandb_project, config=self.config)
36 self._set_steps()
38 def log_metrics(self, metrics, head="train", commit=True):
39 if self.log_wandb:
40 if head:
41 metrics = self._add_head_to_metrics(metrics, head)
42 self._wandb.log(metrics, commit=commit)
43 else:
44 self._wandb.log(metrics, commit=commit)
46 def log_eval_metrics(self, metrics, head="eval"):
47 if self.log_wandb:
48 metrics = self._add_head_to_metrics(metrics, head)
49 for k, v in metrics.items():
50 self._wandb.run.summary[k] = v
52 def _set_steps(self):
53 self._wandb.define_metric("train/*", step_metric="train_step")
54 self._wandb.define_metric("valid/*", step_metric="valid_step")
56 def _add_head_to_metrics(self, metrics, head):
57 head_metrics = dict()
58 for k, v in metrics.items():
59 if "_step" in k:
60 head_metrics[k] = v
61 else:
62 head_metrics[f"{head}/{k}"] = v
64 return head_metrics