matilda-jev-fp4 / kev /tracking.py
yue-maincode's picture
Upload validated MATILDA JEV FP4 model and Decision Index scores
c69aaec verified
Raw History Blame Contribute Delete
3.59 kB
"""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)