File size: 6,980 Bytes
707d9ee
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
"""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)


# TRL prints metrics as a python-dict literal on a single line like:
# {'loss': 3.19, 'grad_norm': 1.45, ...}
_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`.
    """
    # PowerShell `>` redirection writes UTF-16-LE with BOM. Detect via BOM.
    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
    # De-dup: HF Jobs replays log chunks, so we see each step multiple times.
    # Identity is (epoch, loss) — a (epoch, loss) pair is unique per step
    # within a phase.
    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]:
        # try the last row in case keys differ
        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")

    # README expects these filenames specifically:
    # outputs/plots/reward_curve.png
    # outputs/plots/loss_curve.png
    # We emit those as aliases of the most-relevant single-panel plot.
    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))