93 lines
2.9 KiB
Python
93 lines
2.9 KiB
Python
|
|
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
|