lsh9034's picture
Add files using upload-large-folder tool
76d61a0 verified
Raw History Blame Contribute Delete
2.17 kB
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()