0528_deep_hw / page_03_plot_curve.py
yb1n's picture
Upload folder using huggingface_hub
845426e verified
Raw
History Blame Contribute Delete
1.95 kB
import argparse
import json
from pathlib import Path
import matplotlib.pyplot as plt
from trainer_utils import get_best_checkpoint
def parse_args():
parser = argparse.ArgumentParser(description="Plot learning_curve.jsonl")
parser.add_argument("--run-dir", required=True)
return parser.parse_args()
def load_curve(path):
rows = []
with Path(path).open(encoding="utf-8") as file:
for line in file:
rows.append(json.loads(line))
by_step = {}
for row in rows:
step = row.get("step")
if step is not None:
by_step[step] = {**by_step.get(step, {}), **row}
return [by_step[step] for step in sorted(by_step)]
def plot_metric(axis, rows, key, label, color):
selected = [row for row in rows if key in row and row.get("step") is not None]
if selected:
axis.plot(
[row["step"] for row in selected],
[row[key] for row in selected],
label=label,
color=color,
)
axis.set_ylabel(label)
axis.legend()
axis.grid(True)
def main():
args = parse_args()
run_dir = Path(args.run_dir)
rows = load_curve(run_dir / "learning_curve.jsonl")
_, best_step = get_best_checkpoint(run_dir)
fig, axes = plt.subplots(4, 1, figsize=(10, 12), sharex=True)
plot_metric(axes[0], rows, "loss", "train_loss", "tab:blue")
plot_metric(axes[1], rows, "eval_accuracy", "eval_accuracy", "tab:orange")
plot_metric(axes[2], rows, "eval_precision", "eval_precision", "tab:green")
plot_metric(axes[3], rows, "eval_f1", "eval_f1", "tab:red")
if best_step is not None:
for axis in axes:
axis.axvline(best_step, linestyle="--", color="black", label="best")
axes[-1].set_xlabel("step")
fig.tight_layout()
out_path = run_dir / "learning_curve.png"
fig.savefig(out_path, dpi=150)
print("saved:", out_path)
if __name__ == "__main__":
main()