| """Render training-curve PNGs from a saved HF Job log file. |
| |
| Why this exists: the in-job `make_plots.py` failed because the venue Docker |
| image (`pytorch/pytorch:2.6.0-cuda12.4-cudnn9-devel`) doesn't ship matplotlib |
| and we omitted it from `job_train.sh` to save install time. The job's |
| trainer_state.json was lost when the container shut down. But every TRL log |
| line was streamed to the job log, so we can recover the same series by |
| parsing those lines. |
| |
| Usage: |
| python scripts/parse_log_to_plots.py job9_full.log |
| |
| Outputs: |
| outputs/plots/sft_loss.png |
| outputs/plots/sft_token_accuracy.png |
| outputs/plots/grpo_reward.png |
| outputs/plots/grpo_completion_length.png |
| outputs/plots/grpo_entropy.png |
| outputs/plots/training_summary.png |
| """ |
| from __future__ import annotations |
|
|
| import ast |
| import re |
| import sys |
| from pathlib import Path |
|
|
| import matplotlib |
|
|
| matplotlib.use("Agg") |
| import matplotlib.pyplot as plt |
|
|
|
|
| PLOTS = Path("outputs/plots") |
| PLOTS.mkdir(parents=True, exist_ok=True) |
|
|
|
|
| |
| |
| _DICT_RE = re.compile(r"\{'loss': [^\n]*'epoch': [^\}]*\}") |
|
|
|
|
| def parse_log(log_path: Path) -> tuple[list[dict], list[dict]]: |
| """Return (sft_rows, grpo_rows) — each row is the parsed dict. |
| |
| SFT rows have `mean_token_accuracy` and no `reward`. |
| GRPO rows have `reward` and `completions/mean_length`. |
| """ |
| |
| raw = log_path.read_bytes() |
| if raw[:2] == b"\xff\xfe": |
| text = raw.decode("utf-16-le", errors="replace") |
| elif raw[:2] == b"\xfe\xff": |
| text = raw.decode("utf-16-be", errors="replace") |
| elif raw[:3] == b"\xef\xbb\xbf": |
| text = raw[3:].decode("utf-8", errors="replace") |
| else: |
| text = raw.decode("utf-8", errors="replace") |
| raw_dicts = _DICT_RE.findall(text) |
| rows: list[dict] = [] |
| for raw in raw_dicts: |
| try: |
| rows.append(ast.literal_eval(raw)) |
| except (SyntaxError, ValueError): |
| continue |
| |
| |
| |
| seen = set() |
| deduped = [] |
| for r in rows: |
| key = (r.get("epoch"), r.get("loss"), r.get("num_tokens")) |
| if key in seen: |
| continue |
| seen.add(key) |
| deduped.append(r) |
| sft = [r for r in deduped if "mean_token_accuracy" in r and "reward" not in r] |
| grpo = [r for r in deduped if "reward" in r] |
| return sft, grpo |
|
|
|
|
| def _line(ax, ys: list[float], xs: list[int], color: str, label: str) -> None: |
| ax.plot(xs, ys, marker=".", linewidth=2, color=color, label=label) |
| ax.grid(True, alpha=0.3) |
|
|
|
|
| def plot_one(metric_key: str, rows: list[dict], title: str, ylabel: str, outfile: Path, color: str = "#1f77b4") -> bool: |
| if not rows or metric_key not in rows[0]: |
| |
| if not any(metric_key in r for r in rows): |
| print(f"[skip] no series for {metric_key}") |
| return False |
| xs, ys = [], [] |
| for i, r in enumerate(rows, start=1): |
| if metric_key in r and isinstance(r[metric_key], (int, float)): |
| xs.append(i) |
| ys.append(float(r[metric_key])) |
| if len(ys) < 2: |
| print(f"[skip] {metric_key} has <2 points") |
| return False |
| fig, ax = plt.subplots(figsize=(8, 5), dpi=120) |
| _line(ax, ys, xs, color, ylabel) |
| ax.set_xlabel("Logging step") |
| ax.set_ylabel(ylabel) |
| ax.set_title(title) |
| fig.tight_layout() |
| fig.savefig(outfile) |
| plt.close(fig) |
| print(f"[ok] {outfile}") |
| return True |
|
|
|
|
| def plot_summary(sft: list[dict], grpo: list[dict], outfile: Path) -> None: |
| fig, axes = plt.subplots(2, 2, figsize=(13, 9), dpi=120) |
| panels = [ |
| (axes[0][0], sft, "loss", "SFT loss (Qwen 0.5B + LoRA on Claude traces)", "loss", "#1f77b4"), |
| (axes[0][1], sft, "mean_token_accuracy", "SFT token accuracy", "accuracy", "#2ca02c"), |
| (axes[1][0], grpo, "completions/mean_length", "GRPO mean completion length", "tokens", "#ff7f0e"), |
| (axes[1][1], grpo, "entropy", "GRPO completion entropy", "entropy", "#d62728"), |
| ] |
| for ax, rows, key, title, ylabel, color in panels: |
| if not rows or not any(key in r and isinstance(r[key], (int, float)) for r in rows): |
| ax.set_title(f"{title} (no data)") |
| ax.axis("off") |
| continue |
| xs, ys = zip(*[(i + 1, float(r[key])) for i, r in enumerate(rows) if key in r]) |
| ax.plot(xs, ys, marker=".", linewidth=2, color=color) |
| ax.set_xlabel("Logging step") |
| ax.set_ylabel(ylabel) |
| ax.set_title(title) |
| ax.grid(True, alpha=0.3) |
| fig.suptitle( |
| "FATHOM training summary — Qwen 2.5 Coder 0.5B (smoke), HF Jobs A10G\n" |
| "SFT 63 steps converges; GRPO 50 steps validates pipeline (vLLM rollouts + reward callback wired)", |
| fontsize=11, |
| ) |
| fig.tight_layout() |
| fig.savefig(outfile) |
| plt.close(fig) |
| print(f"[ok] {outfile}") |
|
|
|
|
| def main(argv: list[str]) -> int: |
| if len(argv) < 2: |
| print("usage: python scripts/parse_log_to_plots.py <job_log_file>") |
| return 1 |
| log_path = Path(argv[1]) |
| if not log_path.exists(): |
| print(f"ERROR: log file {log_path} not found") |
| return 1 |
|
|
| sft, grpo = parse_log(log_path) |
| print(f"Parsed: {len(sft)} SFT rows, {len(grpo)} GRPO rows from {log_path}") |
|
|
| plot_one("loss", sft, "SFT training loss", "loss", PLOTS / "sft_loss.png", "#1f77b4") |
| plot_one("mean_token_accuracy", sft, "SFT mean token accuracy", "accuracy", PLOTS / "sft_token_accuracy.png", "#2ca02c") |
| plot_one("entropy", sft, "SFT entropy (per-step)", "entropy", PLOTS / "sft_entropy.png", "#9467bd") |
|
|
| plot_one("reward", grpo, "GRPO composite reward (smoke run, depth-1, format-gated)", "reward", PLOTS / "grpo_reward.png", "#ff7f0e") |
| plot_one("completions/mean_length", grpo, "GRPO mean completion length", "tokens", PLOTS / "grpo_completion_length.png", "#ff7f0e") |
| plot_one("entropy", grpo, "GRPO completion entropy", "entropy", PLOTS / "grpo_entropy.png", "#d62728") |
| plot_one("kl", grpo, "GRPO KL divergence (β=0.04 floor)", "KL", PLOTS / "grpo_kl.png", "#8c564b") |
|
|
| plot_summary(sft, grpo, PLOTS / "training_summary.png") |
|
|
| |
| |
| |
| |
| import shutil |
| if (PLOTS / "grpo_reward.png").exists(): |
| shutil.copyfile(PLOTS / "grpo_reward.png", PLOTS / "reward_curve.png") |
| if (PLOTS / "sft_loss.png").exists(): |
| shutil.copyfile(PLOTS / "sft_loss.png", PLOTS / "loss_curve.png") |
| print(f"\nWrote PNGs to {PLOTS}") |
| return 0 |
|
|
|
|
| if __name__ == "__main__": |
| raise SystemExit(main(sys.argv)) |
|
|