batch-size-experiments / plot_run_metrics.py
3v324v23's picture
Add experiment tooling and training analysis
7cffee1
Raw History Blame Contribute Delete
2.43 kB
import argparse
import csv
from collections import defaultdict
from itertools import cycle
from pathlib import Path
import matplotlib.pyplot as plt
COLORS = ["#0072B2", "#D55E00", "#009E73"]
def read_metrics(path: Path) -> dict[str, list[dict[str, float]]]:
metrics = defaultdict(list)
with path.open(encoding="utf-8", newline="") as metrics_file:
for row in csv.DictReader(metrics_file):
metrics[row["run"]].append(
{
"tokens": float(row["tokens"]),
"loss": float(row["loss"]),
"wall_time_seconds": float(row["wall_time_seconds"]),
}
)
if not metrics:
raise ValueError(f"No metric rows found in {path}")
return dict(metrics)
def plot_metrics(
metrics: dict[str, list[dict[str, float]]],
x_field: str,
x_scale: float,
x_label: str,
output_path: Path,
) -> None:
figure, axis = plt.subplots(figsize=(10, 6))
for color, (run, rows) in zip(cycle(COLORS), metrics.items()):
x_values = [row[x_field] / x_scale for row in rows]
loss_values = [row["loss"] for row in rows]
axis.plot(x_values, loss_values, color=color, linewidth=1.8, label=run)
axis.set_xlabel(x_label)
axis.set_ylabel("Training loss")
axis.set_title(f"Training loss vs {x_label.lower()}")
axis.grid(True, alpha=0.25)
axis.legend(title="Run", fontsize=8)
figure.tight_layout()
figure.savefig(output_path, dpi=160)
plt.close(figure)
def main() -> None:
argument_parser = argparse.ArgumentParser()
argument_parser.add_argument("metrics", type=Path)
argument_parser.add_argument("--output-dir", type=Path, required=True)
args = argument_parser.parse_args()
metrics = read_metrics(args.metrics)
args.output_dir.mkdir(parents=True, exist_ok=True)
plot_metrics(
metrics,
x_field="wall_time_seconds",
x_scale=3600,
x_label="Wall-clock time (hours)",
output_path=args.output_dir / "loss_vs_time.png",
)
plot_metrics(
metrics,
x_field="tokens",
x_scale=1_000_000,
x_label="Tokens processed (millions)",
output_path=args.output_dir / "loss_vs_tokens.png",
)
print(f"Wrote {args.output_dir / 'loss_vs_time.png'}")
print(f"Wrote {args.output_dir / 'loss_vs_tokens.png'}")
if __name__ == "__main__":
main()