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()