File size: 2,167 Bytes
76d61a0
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
from __future__ import annotations

from pathlib import Path
from typing import Any


class ExperimentLogger:
    def __init__(self, config: dict[str, Any], mode: str):
        self.config = config
        self.mode = str(mode)
        self.enabled = bool(config.get("logging", {}).get("wandb", {}).get("enabled", False))
        self._wandb = None
        self._run = None

    def start(self) -> None:
        if not self.enabled:
            return
        try:
            import wandb
        except ImportError as exc:
            raise ImportError("wandb logging is enabled, but wandb is not installed") from exc

        wandb_cfg = dict(self.config.get("logging", {}).get("wandb", {}))
        project = str(wandb_cfg.get("project", "atmos_project"))
        run_name = wandb_cfg.get("run_name") or Path(str(self.config.get("output_dir", "run"))).name
        if self.mode != "train":
            run_name = f"{run_name}_{self.mode}"
        self._wandb = wandb
        self._run = wandb.init(
            project=project,
            name=run_name,
            config={k: v for k, v in self.config.items() if not str(k).startswith("_")},
            tags=wandb_cfg.get("tags"),
            notes=wandb_cfg.get("notes"),
            mode=str(wandb_cfg.get("mode", "online")),
        )

    def log(self, metrics: dict[str, Any], step: int | None = None, prefix: str | None = None) -> None:
        if not self.enabled or self._wandb is None:
            return
        payload = {}
        for key, value in metrics.items():
            if isinstance(value, (int, float, str, bool)) or value is None:
                payload[f"{prefix}/{key}" if prefix else str(key)] = value
        self._wandb.log(payload, step=step)

    def log_file(self, path: str | Path, name: str | None = None) -> None:
        if not self.enabled or self._wandb is None:
            return
        artifact = self._wandb.Artifact(name or Path(path).stem, type=f"{self.mode}_result")
        artifact.add_file(str(path))
        self._wandb.log_artifact(artifact)

    def finish(self) -> None:
        if self.enabled and self._wandb is not None:
            self._wandb.finish()