File size: 3,593 Bytes
c69aaec
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
"""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)