Download plot_run_metrics.py from adamskrodzki/batch-size-experiments: direct link, hf CLI and curl.
- Browser
- Download file 2.43 kB
-
https://huggingface.co/adamskrodzki/batch-size-experiments/resolve/main/plot_run_metrics.py
- Command line
-
hf download hf://adamskrodzki/batch-size-experiments/plot_run_metrics.py
-
curl -L -o plot_run_metrics.py https://huggingface.co/adamskrodzki/batch-size-experiments/resolve/main/plot_run_metrics.py
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() | |