Download code/training/src/training_validation/logger.py from lsh9034/ci-net: direct link, hf CLI and curl.
- Browser
- Download file 2.17 kB
-
https://huggingface.co/lsh9034/ci-net/resolve/main/code/training/src/training_validation/logger.py
- Command line
-
hf download hf://lsh9034/ci-net/code/training/src/training_validation/logger.py
-
curl -L -o logger.py https://huggingface.co/lsh9034/ci-net/resolve/main/code/training/src/training_validation/logger.py
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() | |