2026-07-25 08:27:22 +00:00
|
|
|
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
|
2026-07-25 08:27:22 +00:00
|
|
|
|
|
|
|
|
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
|
|
|
|
2026-07-25 08:27:22 +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
|
|
|
|
2026-07-25 08:27:22 +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,
|
2026-07-25 08:27:22 +00:00
|
|
|
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"
|
|
|
|
|
|
|
|
|
|
|
2026-07-25 08:27:22 +00:00
|
|
|
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
|