airfRANS-model-exploration/src/airfrans_frontier/training/observability.py

93 lines
2.9 KiB
Python
Raw Normal View History

from __future__ import annotations
from dataclasses import asdict, is_dataclass
from pathlib import Path
from typing import Any, Mapping
from airfrans_frontier.training.config import TrainingConfig
class TrainingObserver:
@property
def url(self) -> str | None:
return None
def log(self, metrics: Mapping[str, Any]) -> None:
return None
def update_summary(self, metrics: Mapping[str, Any]) -> None:
return None
def finish(self, *, exit_code: int = 0) -> None:
return None
class WandbObserver(TrainingObserver):
def __init__(self, run: Any, wandb_module: Any) -> None:
self._run = run
self._wandb = wandb_module
@property
def url(self) -> str | None:
get_url = getattr(self._run, "get_url", None)
if callable(get_url):
return get_url()
url = getattr(self._run, "url", None)
return str(url) if url else None
def log(self, metrics: Mapping[str, Any]) -> None:
payload = _json_safe(dict(metrics))
step = payload.get("step")
if isinstance(step, int):
self._wandb.log(payload, step=step)
else:
self._wandb.log(payload)
def update_summary(self, metrics: Mapping[str, Any]) -> None:
for key, value in _json_safe(dict(metrics)).items():
self._run.summary[key] = value
def finish(self, *, exit_code: int = 0) -> None:
self._wandb.finish(exit_code=exit_code)
def start_observer(config: TrainingConfig, *, run_dir: Path) -> TrainingObserver:
observability = config.observability
if observability.backend == "none" or observability.mode == "disabled":
return TrainingObserver()
if observability.backend != "wandb":
raise ValueError(f"Unsupported observability backend: {observability.backend}")
try:
import wandb
except ModuleNotFoundError as exc:
raise RuntimeError("wandb is required when [observability].backend = 'wandb'") from exc
wandb_dir = run_dir.parent / ".wandb"
run = wandb.init(
entity=observability.entity,
project=observability.project,
name=config.run.name,
tags=list(observability.tags),
mode=observability.mode,
config=_json_safe(asdict(config)),
dir=str(wandb_dir),
)
if run is not None:
run.define_metric("step")
run.define_metric("*", step_metric="step")
return WandbObserver(run, wandb)
def _json_safe(value: Any) -> Any:
if is_dataclass(value):
return _json_safe(asdict(value))
if isinstance(value, Path):
return str(value)
if isinstance(value, Mapping):
return {str(key): _json_safe(item) for key, item in value.items()}
if isinstance(value, tuple):
return [_json_safe(item) for item in value]
if isinstance(value, list):
return [_json_safe(item) for item in value]
return value