"""Optional Weights & Biases tracking for the trainers (rank 0 only). Disabled unless a project is given. The run id is stored in the run directory so `--resume` continues the same W&B run. Tracking never stops training: if W&B fails to start or log, a warning is printed and training continues. """ import sys import uuid from collections.abc import Mapping from pathlib import Path from typing import Any, cast from kev.evaluate import Metrics SCALARS = ("accuracy", "ece", "brier", "nll", "score_mae", "soft_nll") def flatten_metrics(prefix: str, values: Mapping[str, Any] | Metrics) -> dict[str, float]: return {f"{prefix}/{name}": float(values[name]) for name in SCALARS if name in values} # type: ignore[literal-required] class Tracker: def __init__(self, run: Path, project: str | None, mode: str, config: Mapping[str, object], enabled: bool = True) -> None: self.wandb: Any = None if not enabled or not project or mode == "disabled": return try: import wandb identity = run / "wandb-id.txt" run_id = identity.read_text().strip() if identity.exists() else uuid.uuid4().hex[:12] identity.write_text(run_id + "\n") directory = run / "wandb" directory.mkdir(parents=True, exist_ok=True) wandb.init(project=project, name=run.name, id=run_id, resume="allow", mode=cast(Any, mode), dir=str(directory), config=dict(config)) self.wandb = wandb except Exception as error: # noqa: BLE001 - tracking must never stop training print(f"WARNING: W&B disabled ({type(error).__name__}: {error})", file=sys.stderr, flush=True) def log(self, values: Mapping[str, float], step: int) -> None: if self.wandb is None: return try: self.wandb.log(dict(values), step=step) except Exception as error: # noqa: BLE001 print(f"WARNING: W&B log failed at step {step}: {error}", file=sys.stderr, flush=True) def log_training(self, value: Mapping[str, Any]) -> None: keys = ("loss", "learning_rate", "gradient_norm", "tokens_per_second", "step_seconds", "gpu_peak_gb", "input_tokens", "examples_seen") self.log({f"train/{key}": float(value[key]) for key in keys if value.get(key) is not None}, int(value["step"])) def log_evaluation(self, step: int, raw: Metrics, fitted: Metrics, panels: Mapping[str, Metrics], temperature: float, seconds: float) -> None: values = {**flatten_metrics("dev", fitted), **flatten_metrics("dev_raw", raw), "dev/temperature": temperature, "dev/evaluation_seconds": seconds} for name, metrics in panels.items(): values.update(flatten_metrics(f"panel/{name}", metrics)) self.log(values, step) def summary(self, values: Mapping[str, object]) -> None: if self.wandb is None: return try: for key, value in values.items(): if isinstance(value, (int, float, str)) and not isinstance(value, bool): self.wandb.run.summary[key] = value except Exception as error: # noqa: BLE001 print(f"WARNING: W&B summary failed: {error}", file=sys.stderr, flush=True) def finish(self) -> None: if self.wandb is not None: try: self.wandb.finish() except Exception as error: # noqa: BLE001 print(f"WARNING: W&B finish failed: {error}", file=sys.stderr, flush=True)