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