from __future__ import annotations from dataclasses import asdict, is_dataclass from pathlib import Path 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 def update_config(self, values: Mapping[str, Any]) -> None: return None def log_artifact_files( self, run_dir: Path, names: Iterable[str], *, event: str, step: int, aliases: Iterable[str] = (), ) -> 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 update_config(self, values: Mapping[str, Any]) -> None: self._run.config.update(_json_safe(dict(values)), allow_val_change=True) 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}"]) 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, 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) 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