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

141 lines
4.5 KiB
Python
Raw Normal View History

from __future__ import annotations
from dataclasses import asdict, is_dataclass
from pathlib import Path
2026-07-27 17:51:28 +00:00
from typing import Any, Iterable, 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
2026-07-25 16:12:49 +00:00
def update_config(self, values: Mapping[str, Any]) -> None:
return None
2026-07-27 17:51:28 +00:00
def log_artifact_files(
self,
run_dir: Path,
names: Iterable[str],
*,
event: str,
step: int,
aliases: Iterable[str] = (),
) -> None:
return None
2026-07-25 16:12:49 +00:00
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
2026-07-25 16:12:49 +00:00
def update_config(self, values: Mapping[str, Any]) -> None:
self._run.config.update(_json_safe(dict(values)), allow_val_change=True)
2026-07-27 17:51:28 +00:00
def log_artifact_files(
self,
run_dir: Path,
names: Iterable[str],
*,
event: str,
step: int,
aliases: Iterable[str] = (),
) -> None:
existing = [name for name in names if (run_dir / name).is_file()]
if not existing:
return
artifact_name = _artifact_name(getattr(self._run, "name", "airfrans-run"), event, step)
artifact = self._wandb.Artifact(
artifact_name,
type="airfrans-run-artifacts",
metadata={
"event": event,
"step": step,
"file_count": len(existing),
},
)
for name in existing:
artifact.add_file(str(run_dir / name), name=name)
self._run.log_artifact(artifact, aliases=list(aliases) or [event, f"step-{step}"])
2026-07-25 16:12:49 +00:00
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,
2026-07-25 16:12:49 +00:00
group=observability.group,
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)
2026-07-27 17:51:28 +00:00
def _artifact_name(run_name: Any, event: str, step: int) -> str:
raw = f"{run_name}-{event}-step-{step}"
safe = "".join(character if character.isalnum() or character in "-_." else "-" for character in str(raw))
return safe.strip("-") or "airfrans-run-artifacts"
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