analyze_training

Turn a marola-sea training run into a report you can act on — and hand back to an agent.

Both trainers set report_to=[] (finetune/train_lora.py, finetune/train_dpo.py), so there is no MLflow or W&B run to open: the evidence is whatever the job left on disk. That is more than the scrollback, though. Hugging Face's Trainer writes trainer_state.json next to the adapter and inside every checkpoint-N/, and its log_history is the whole run as structured records — every logged train loss, every eval, and the final throughput summary. This reads that (falling back to scraping stdout when a state file is missing) and answers the question the raw numbers do not: what should the next run change?

python3 scripts/analyze_training.py ../marola-checkpoints/tiny/adapter
python3 scripts/analyze_training.py <dir-or-state.json> ... --markdown report.md --json report.json

The findings are deliberately phrased as a config diff — num_train_epochs, learning_rate, beta — because that is what the next run actually needs, and what is useful to paste back into a session as "here is what happened, what do we change?".

  1#!/usr/bin/env python3
  2"""Turn a marola-sea training run into a report you can act on — and hand back to an agent.
  3
  4Both trainers set `report_to=[]` (finetune/train_lora.py, finetune/train_dpo.py), so there is no
  5MLflow or W&B run to open: the evidence is whatever the job left on disk. That is more than the
  6scrollback, though. Hugging Face's Trainer writes `trainer_state.json` next to the adapter and
  7inside every `checkpoint-N/`, and its `log_history` is the whole run as structured records —
  8every logged train loss, every eval, and the final throughput summary. This reads that (falling
  9back to scraping stdout when a state file is missing) and answers the question the raw numbers do
 10not: *what should the next run change?*
 11
 12    python3 scripts/analyze_training.py ../marola-checkpoints/tiny/adapter
 13    python3 scripts/analyze_training.py <dir-or-state.json> ... --markdown report.md --json report.json
 14
 15The findings are deliberately phrased as a config diff — `num_train_epochs`, `learning_rate`,
 16`beta` — because that is what the next run actually needs, and what is useful to paste back into a
 17session as "here is what happened, what do we change?".
 18"""
 19
 20from __future__ import annotations
 21
 22import argparse
 23import json
 24import re
 25import sys
 26from pathlib import Path
 27
 28# HF logs a dict per logging_steps; DPO adds the reward keys trl's DPOTrainer emits.
 29TRAIN_KEY, EVAL_KEY = "loss", "eval_loss"
 30# How much the loss must still be falling across the last logged steps to call a run
 31# undertrained. A couple of percent is ordinary step-to-step noise on a converged curve.
 32STILL_FALLING = 0.05
 33DPO_ACC_KEY = "rewards/accuracies"
 34
 35
 36def find_state(target: Path) -> Path | None:
 37    """trainer_state.json for a run: in the output dir, else the newest checkpoint-N inside it."""
 38    if target.is_file():
 39        return target
 40    direct = target / "trainer_state.json"
 41    if direct.is_file():
 42        return direct
 43    checkpoints = sorted(
 44        (p for p in target.glob("checkpoint-*") if (p / "trainer_state.json").is_file()),
 45        key=lambda p: int(p.name.split("-")[-1]),
 46    )
 47    return checkpoints[-1] / "trainer_state.json" if checkpoints else None
 48
 49
 50def parse_stdout(text: str) -> list[dict]:
 51    """Fallback for a run whose state file is gone: HF prints each log dict with single quotes."""
 52    records = []
 53    for match in re.finditer(r"\{'(?:loss|eval_loss|train_runtime)'.*?\}", text):
 54        try:
 55            records.append(json.loads(match.group(0).replace("'", '"')))
 56        except json.JSONDecodeError:
 57            continue
 58    return records
 59
 60
 61def load_run(target: Path) -> dict:
 62    state = find_state(target)
 63    if state is not None:
 64        data = json.loads(state.read_text())
 65        return {
 66            "name": target.name if target.is_dir() else target.parent.name,
 67            "source": str(state),
 68            "log_history": data.get("log_history", []),
 69            "epochs_run": data.get("epoch"),
 70            "global_step": data.get("global_step"),
 71        }
 72    logs = sorted(target.glob("*.log")) if target.is_dir() else []
 73    if logs:
 74        history = [r for f in logs for r in parse_stdout(f.read_text(errors="replace"))]
 75        return {
 76            "name": target.name,
 77            "source": f"{logs[0]} (stdout)",
 78            "log_history": history,
 79            "epochs_run": None,
 80            "global_step": None,
 81        }
 82    raise FileNotFoundError(f"no trainer_state.json or *.log under {target}")
 83
 84
 85def series(history: list[dict], key: str) -> list[tuple[float, float]]:
 86    """(epoch, value) pairs for one metric, in order, skipping records that lack it."""
 87    return [(r.get("epoch", i), r[key]) for i, r in enumerate(history) if key in r]
 88
 89
 90def summarize(run: dict) -> dict:
 91    history = run["log_history"]
 92    train, ev = series(history, TRAIN_KEY), series(history, EVAL_KEY)
 93    final = next((r for r in reversed(history) if "train_runtime" in r), {})
 94    dpo = series(history, DPO_ACC_KEY)
 95    return {
 96        "name": run["name"],
 97        "source": run["source"],
 98        "kind": "dpo" if dpo else "sft",
 99        "steps": run.get("global_step") or len(train),
100        "epochs_run": run.get("epochs_run"),
101        "train_loss": [v for _, v in train],
102        "eval_loss": [v for _, v in ev],
103        "eval_epochs": [e for e, _ in ev],
104        "grad_norm": [v for _, v in series(history, "grad_norm")],
105        "dpo_accuracy": [v for _, v in dpo],
106        "runtime_s": final.get("train_runtime"),
107        "samples_per_s": final.get("train_samples_per_second"),
108    }
109
110
111def _trend(values: list[float], tail: int = 3) -> float:
112    """How much the metric moved across its last `tail` points, as a fraction of its own scale."""
113    if len(values) < tail + 1:
114        return 0.0
115    recent, earlier = values[-tail:], values[-(tail + 1)]
116    return (earlier - sum(recent) / len(recent)) / abs(earlier) if earlier else 0.0
117
118
119def diagnose(s: dict) -> list[dict]:
120    """Findings as {severity, title, evidence, suggestion} — suggestion names a config knob."""
121    out: list[dict] = []
122    train, ev, steps = s["train_loss"], s["eval_loss"], s["steps"]
123
124    # The failure a small corpus actually hits: too few optimizer steps to learn anything.
125    if steps and steps < 50:
126        out.append(
127            {
128                "severity": "high",
129                "title": "the run was too short to learn much",
130                "evidence": f"{steps} optimizer steps in total",
131                "suggestion": "raise num_train_epochs, or lower gradient_accumulation_steps "
132                "(SFT uses 4, DPO 2) so the same data yields more steps; a few dozen "
133                "steps mostly measures the initialisation, not the dataset",
134            }
135        )
136
137    overfit = len(ev) >= 2 and ev[-1] > min(ev) * 1.02
138    if overfit:
139        best = ev.index(min(ev))
140        out.append(
141            {
142                "severity": "high",
143                "title": "eval loss turned back up — overfitting",
144                "evidence": f"best eval_loss {min(ev):.4f} at epoch "
145                f"{s['eval_epochs'][best]:.0f}, ended at {ev[-1]:.4f}",
146                "suggestion": f"train for ~{s['eval_epochs'][best]:.0f} epochs instead, or add data; "
147                "the later epochs made the model worse on held-out examples",
148            }
149        )
150
151    if len(train) >= 4:
152        drop = (train[0] - train[-1]) / train[0] if train[0] else 0.0
153        if drop < 0.05:
154            out.append(
155                {
156                    "severity": "high",
157                    "title": "train loss barely moved",
158                    "evidence": f"{train[0]:.4f} -> {train[-1]:.4f} ({drop * 100:.1f}%)",
159                    "suggestion": "raise learning_rate (SFT default 2e-4 for LoRA), or check the "
160                    "dataset actually loaded — a flat curve from step 0 is usually one "
161                    "of those two, not a model problem",
162                }
163            )
164        # Not `elif overfit`: a train loss still falling while eval rises IS the
165        # overfitting above, and telling you to train longer would contradict it.
166        elif _trend(train) > STILL_FALLING and not overfit:
167            out.append(
168                {
169                    "severity": "info",
170                    "title": "still improving when it stopped — undertrained",
171                    "evidence": f"loss fell {_trend(train) * 100:.1f}% across the last logged steps",
172                    "suggestion": "raise num_train_epochs; the curve had not flattened yet",
173                }
174            )
175
176    if s["grad_norm"]:
177        peak, median = max(s["grad_norm"]), sorted(s["grad_norm"])[len(s["grad_norm"]) // 2]
178        if median and peak > median * 10:
179            out.append(
180                {
181                    "severity": "medium",
182                    "title": "gradient-norm spikes",
183                    "evidence": f"peak {peak:.2f} against a median of {median:.2f}",
184                    "suggestion": "lower learning_rate or add warmup_ratio=0.03; spikes this size "
185                    "mean some steps moved the weights far more than the rest",
186                }
187            )
188
189    if s["kind"] == "dpo" and s["dpo_accuracy"]:
190        final_acc = sum(s["dpo_accuracy"][-3:]) / len(s["dpo_accuracy"][-3:])
191        if final_acc < 0.6:
192            out.append(
193                {
194                    "severity": "high",
195                    "title": "DPO is not separating chosen from rejected",
196                    "evidence": f"rewards/accuracies averaged {final_acc:.2f} at the end (0.5 = chance)",
197                    "suggestion": "raise beta (default 0.1) so the preference signal counts for more, "
198                    "or check build_dpo_dataset.py — pairs that are near-identical give "
199                    "the trainer nothing to separate",
200                }
201            )
202
203    if not out:
204        out.append(
205            {
206                "severity": "info",
207                "title": "nothing anomalous in the curves",
208                "evidence": "loss fell and eval did not diverge",
209                "suggestion": "scale up: the next question is the preset, not the schedule",
210            }
211        )
212    return out
213
214
215def render_markdown(runs: list[tuple[dict, list[dict]]]) -> str:
216    rank = {"high": "🔴", "medium": "🟡", "info": "🔵"}
217    lines = ["# marola-sea training report", ""]
218    for s, findings in runs:
219        lines += [
220            f"## {s['name']} ({s['kind'].upper()})",
221            "",
222            f"- source: `{s['source']}`",
223            f"- steps: {s['steps']}, epochs run: {s['epochs_run'] or 'n/a'}",
224        ]
225        if s["train_loss"]:
226            lines.append(f"- train loss: {s['train_loss'][0]:.4f} → {s['train_loss'][-1]:.4f}")
227        if s["eval_loss"]:
228            lines.append(
229                f"- eval loss: {s['eval_loss'][0]:.4f} → {s['eval_loss'][-1]:.4f} "
230                f"(best {min(s['eval_loss']):.4f})"
231            )
232        if s["dpo_accuracy"]:
233            lines.append(f"- DPO reward accuracy: ends at {s['dpo_accuracy'][-1]:.2f}")
234        if s["runtime_s"]:
235            lines.append(
236                f"- runtime: {s['runtime_s'] / 60:.1f} min"
237                + (f", {s['samples_per_s']:.2f} samples/s" if s["samples_per_s"] else "")
238            )
239        lines += ["", "### What to change next run", ""]
240        for f in findings:
241            lines += [
242                f"**{rank.get(f['severity'], '•')} {f['title']}**  ",
243                f"{f['evidence']}  ",
244                f"→ {f['suggestion']}",
245                "",
246            ]
247    return "\n".join(lines)
248
249
250def self_test() -> int:
251    fails = 0
252
253    def ok(got, want, what):
254        nonlocal fails
255        if got == want:
256            print(f"  ok   {what}")
257        else:
258            print(f"  FAIL {what} — got {got!r} want {want!r}")
259            fails += 1
260
261    def titles(summary):
262        return [f["title"] for f in diagnose(summary)]
263
264    base = {
265        "name": "t",
266        "source": "x",
267        "kind": "sft",
268        "steps": 400,
269        "epochs_run": 3,
270        "train_loss": [],
271        "eval_loss": [],
272        "eval_epochs": [],
273        "grad_norm": [],
274        "dpo_accuracy": [],
275        "runtime_s": None,
276        "samples_per_s": None,
277    }
278
279    ok(
280        titles({**base, "steps": 12}),
281        ["the run was too short to learn much"],
282        "a 12-step run is flagged — the failure a small corpus actually hits",
283    )
284    ok(
285        "the run was too short to learn much" in titles({**base, "steps": 400}),
286        False,
287        "a 400-step run is not flagged for length",
288    )
289
290    over = {
291        **base,
292        "eval_loss": [2.0, 1.2, 1.5],
293        "eval_epochs": [1, 2, 3],
294        "train_loss": [3.0, 2.0, 1.0, 0.4],
295    }
296    ok(
297        "eval loss turned back up — overfitting" in titles(over),
298        True,
299        "eval loss rising while train loss falls is called overfitting",
300    )
301    ok(
302        "2" in diagnose(over)[0]["suggestion"],
303        True,
304        "and the suggestion names the epoch the run should have stopped at",
305    )
306
307    ok(
308        "train loss barely moved" in titles({**base, "train_loss": [2.0, 1.99, 1.98, 1.99]}),
309        True,
310        "a flat train curve is flagged as a learning-rate or dataset problem",
311    )
312    ok(
313        "train loss barely moved" in titles({**base, "train_loss": [3.0, 2.0, 1.0, 0.5]}),
314        False,
315        "a curve that fell by half is not",
316    )
317
318    spikes = {**base, "train_loss": [3.0, 2.0, 1.0, 0.5], "grad_norm": [0.5, 0.5, 0.6, 40.0, 0.5]}
319    ok("gradient-norm spikes" in titles(spikes), True, "a 40x gradient spike is flagged")
320    ok(
321        "gradient-norm spikes" in titles({**base, "grad_norm": [0.5, 0.6, 0.55, 0.7]}),
322        False,
323        "a steady gradient norm is not",
324    )
325
326    dpo = {
327        **base,
328        "kind": "dpo",
329        "train_loss": [0.7, 0.6, 0.5, 0.45],
330        "dpo_accuracy": [0.5, 0.52, 0.51],
331    }
332    ok(
333        "DPO is not separating chosen from rejected" in titles(dpo),
334        True,
335        "DPO reward accuracy at chance is flagged — the run learned no preference",
336    )
337    ok(
338        "DPO is not separating chosen from rejected"
339        in titles({**dpo, "dpo_accuracy": [0.8, 0.85, 0.9]}),
340        False,
341        "a DPO run that does separate them is not flagged",
342    )
343
344    ok(
345        titles({**base, "train_loss": [3.0, 2.0, 1.0, 0.5, 0.49, 0.49, 0.49]}),
346        ["nothing anomalous in the curves"],
347        "a healthy converged run reports nothing to change",
348    )
349
350    ok(
351        "still improving when it stopped — undertrained" in titles(over),
352        False,
353        "an overfitting run is never also told to train longer — the advice would contradict",
354    )
355    ok(
356        "still improving when it stopped — undertrained"
357        in titles({**base, "train_loss": [3.0, 2.5, 2.0, 1.5, 1.0]}),
358        True,
359        "but a run that is genuinely still falling, with no eval divergence, is",
360    )
361
362    ok(
363        len(parse_stdout("{'loss': 1.5, 'epoch': 0.5}\nnoise\n{'eval_loss': 1.2, 'epoch': 1.0}")),
364        2,
365        "stdout fallback recovers HF's single-quoted log dicts",
366    )
367    ok(
368        parse_stdout("{'loss': 1.5, 'epoch': 0.5}")[0]["loss"],
369        1.5,
370        "and their values survive the quote conversion",
371    )
372    ok(len(parse_stdout("no log lines here at all")), 0, "prose is not mistaken for a log record")
373
374    hist = [
375        {"loss": 2.0, "epoch": 0.5},
376        {"eval_loss": 1.8, "epoch": 1.0},
377        {"train_runtime": 120.0, "train_samples_per_second": 4.0},
378    ]
379    s = summarize(
380        {"name": "r", "source": "s", "log_history": hist, "epochs_run": 1, "global_step": 10}
381    )
382    ok(
383        (s["train_loss"], s["eval_loss"], s["runtime_s"]),
384        ([2.0], [1.8], 120.0),
385        "summarize splits train, eval and the final throughput record apart",
386    )
387    ok(s["kind"], "sft", "a run with no reward keys is classified as SFT")
388    ok(
389        summarize(
390            {
391                "name": "r",
392                "source": "s",
393                "log_history": [{DPO_ACC_KEY: 0.7, "epoch": 1}],
394                "epochs_run": 1,
395                "global_step": 1,
396            }
397        )["kind"],
398        "dpo",
399        "and one with them as DPO",
400    )
401
402    ok(
403        render_markdown([(s, diagnose(s))]).startswith("# marola-sea training report"),
404        True,
405        "the markdown report renders",
406    )
407
408    print(
409        "analyze_training self-test: ok"
410        if not fails
411        else f"analyze_training self-test: {fails} failure(s)",
412        file=sys.stderr if fails else sys.stdout,
413    )
414    return 1 if fails else 0
415
416
417def main() -> int:
418    ap = argparse.ArgumentParser(description=__doc__.split("\n")[0])
419    ap.add_argument("runs", nargs="*", type=Path, help="adapter dirs, or trainer_state.json files")
420    ap.add_argument("--markdown", type=Path, help="write the report here (default: stdout)")
421    ap.add_argument("--json", type=Path, help="also write the structured summary here")
422    ap.add_argument("--self-test", action="store_true")
423    args = ap.parse_args()
424
425    if args.self_test:
426        return self_test()
427    if not args.runs:
428        ap.error("give at least one run directory (or --self-test)")
429
430    analyzed = []
431    for target in args.runs:
432        try:
433            s = summarize(load_run(target))
434        except (FileNotFoundError, json.JSONDecodeError) as e:
435            print(f"analyze_training: skipping {target}: {e}", file=sys.stderr)
436            continue
437        analyzed.append((s, diagnose(s)))
438    if not analyzed:
439        print("analyze_training: no readable runs", file=sys.stderr)
440        return 1
441
442    report = render_markdown(analyzed)
443    if args.markdown:
444        args.markdown.write_text(report)
445        print(f"wrote {args.markdown}")
446    else:
447        print(report)
448    if args.json:
449        args.json.write_text(
450            json.dumps([{"summary": s, "findings": f} for s, f in analyzed], indent=2)
451        )
452        print(f"wrote {args.json}")
453    return 0
454
455
456if __name__ == "__main__":
457    sys.exit(main())
STILL_FALLING = 0.05
DPO_ACC_KEY = 'rewards/accuracies'
def find_state(target: pathlib.Path) -> pathlib.Path | None:
37def find_state(target: Path) -> Path | None:
38    """trainer_state.json for a run: in the output dir, else the newest checkpoint-N inside it."""
39    if target.is_file():
40        return target
41    direct = target / "trainer_state.json"
42    if direct.is_file():
43        return direct
44    checkpoints = sorted(
45        (p for p in target.glob("checkpoint-*") if (p / "trainer_state.json").is_file()),
46        key=lambda p: int(p.name.split("-")[-1]),
47    )
48    return checkpoints[-1] / "trainer_state.json" if checkpoints else None

trainer_state.json for a run: in the output dir, else the newest checkpoint-N inside it.

def parse_stdout(text: str) -> list[dict]:
51def parse_stdout(text: str) -> list[dict]:
52    """Fallback for a run whose state file is gone: HF prints each log dict with single quotes."""
53    records = []
54    for match in re.finditer(r"\{'(?:loss|eval_loss|train_runtime)'.*?\}", text):
55        try:
56            records.append(json.loads(match.group(0).replace("'", '"')))
57        except json.JSONDecodeError:
58            continue
59    return records

Fallback for a run whose state file is gone: HF prints each log dict with single quotes.

def load_run(target: pathlib.Path) -> dict:
62def load_run(target: Path) -> dict:
63    state = find_state(target)
64    if state is not None:
65        data = json.loads(state.read_text())
66        return {
67            "name": target.name if target.is_dir() else target.parent.name,
68            "source": str(state),
69            "log_history": data.get("log_history", []),
70            "epochs_run": data.get("epoch"),
71            "global_step": data.get("global_step"),
72        }
73    logs = sorted(target.glob("*.log")) if target.is_dir() else []
74    if logs:
75        history = [r for f in logs for r in parse_stdout(f.read_text(errors="replace"))]
76        return {
77            "name": target.name,
78            "source": f"{logs[0]} (stdout)",
79            "log_history": history,
80            "epochs_run": None,
81            "global_step": None,
82        }
83    raise FileNotFoundError(f"no trainer_state.json or *.log under {target}")
def series(history: list[dict], key: str) -> list[tuple[float, float]]:
86def series(history: list[dict], key: str) -> list[tuple[float, float]]:
87    """(epoch, value) pairs for one metric, in order, skipping records that lack it."""
88    return [(r.get("epoch", i), r[key]) for i, r in enumerate(history) if key in r]

(epoch, value) pairs for one metric, in order, skipping records that lack it.

def summarize(run: dict) -> dict:
 91def summarize(run: dict) -> dict:
 92    history = run["log_history"]
 93    train, ev = series(history, TRAIN_KEY), series(history, EVAL_KEY)
 94    final = next((r for r in reversed(history) if "train_runtime" in r), {})
 95    dpo = series(history, DPO_ACC_KEY)
 96    return {
 97        "name": run["name"],
 98        "source": run["source"],
 99        "kind": "dpo" if dpo else "sft",
100        "steps": run.get("global_step") or len(train),
101        "epochs_run": run.get("epochs_run"),
102        "train_loss": [v for _, v in train],
103        "eval_loss": [v for _, v in ev],
104        "eval_epochs": [e for e, _ in ev],
105        "grad_norm": [v for _, v in series(history, "grad_norm")],
106        "dpo_accuracy": [v for _, v in dpo],
107        "runtime_s": final.get("train_runtime"),
108        "samples_per_s": final.get("train_samples_per_second"),
109    }
def diagnose(s: dict) -> list[dict]:
120def diagnose(s: dict) -> list[dict]:
121    """Findings as {severity, title, evidence, suggestion} — suggestion names a config knob."""
122    out: list[dict] = []
123    train, ev, steps = s["train_loss"], s["eval_loss"], s["steps"]
124
125    # The failure a small corpus actually hits: too few optimizer steps to learn anything.
126    if steps and steps < 50:
127        out.append(
128            {
129                "severity": "high",
130                "title": "the run was too short to learn much",
131                "evidence": f"{steps} optimizer steps in total",
132                "suggestion": "raise num_train_epochs, or lower gradient_accumulation_steps "
133                "(SFT uses 4, DPO 2) so the same data yields more steps; a few dozen "
134                "steps mostly measures the initialisation, not the dataset",
135            }
136        )
137
138    overfit = len(ev) >= 2 and ev[-1] > min(ev) * 1.02
139    if overfit:
140        best = ev.index(min(ev))
141        out.append(
142            {
143                "severity": "high",
144                "title": "eval loss turned back up — overfitting",
145                "evidence": f"best eval_loss {min(ev):.4f} at epoch "
146                f"{s['eval_epochs'][best]:.0f}, ended at {ev[-1]:.4f}",
147                "suggestion": f"train for ~{s['eval_epochs'][best]:.0f} epochs instead, or add data; "
148                "the later epochs made the model worse on held-out examples",
149            }
150        )
151
152    if len(train) >= 4:
153        drop = (train[0] - train[-1]) / train[0] if train[0] else 0.0
154        if drop < 0.05:
155            out.append(
156                {
157                    "severity": "high",
158                    "title": "train loss barely moved",
159                    "evidence": f"{train[0]:.4f} -> {train[-1]:.4f} ({drop * 100:.1f}%)",
160                    "suggestion": "raise learning_rate (SFT default 2e-4 for LoRA), or check the "
161                    "dataset actually loaded — a flat curve from step 0 is usually one "
162                    "of those two, not a model problem",
163                }
164            )
165        # Not `elif overfit`: a train loss still falling while eval rises IS the
166        # overfitting above, and telling you to train longer would contradict it.
167        elif _trend(train) > STILL_FALLING and not overfit:
168            out.append(
169                {
170                    "severity": "info",
171                    "title": "still improving when it stopped — undertrained",
172                    "evidence": f"loss fell {_trend(train) * 100:.1f}% across the last logged steps",
173                    "suggestion": "raise num_train_epochs; the curve had not flattened yet",
174                }
175            )
176
177    if s["grad_norm"]:
178        peak, median = max(s["grad_norm"]), sorted(s["grad_norm"])[len(s["grad_norm"]) // 2]
179        if median and peak > median * 10:
180            out.append(
181                {
182                    "severity": "medium",
183                    "title": "gradient-norm spikes",
184                    "evidence": f"peak {peak:.2f} against a median of {median:.2f}",
185                    "suggestion": "lower learning_rate or add warmup_ratio=0.03; spikes this size "
186                    "mean some steps moved the weights far more than the rest",
187                }
188            )
189
190    if s["kind"] == "dpo" and s["dpo_accuracy"]:
191        final_acc = sum(s["dpo_accuracy"][-3:]) / len(s["dpo_accuracy"][-3:])
192        if final_acc < 0.6:
193            out.append(
194                {
195                    "severity": "high",
196                    "title": "DPO is not separating chosen from rejected",
197                    "evidence": f"rewards/accuracies averaged {final_acc:.2f} at the end (0.5 = chance)",
198                    "suggestion": "raise beta (default 0.1) so the preference signal counts for more, "
199                    "or check build_dpo_dataset.py — pairs that are near-identical give "
200                    "the trainer nothing to separate",
201                }
202            )
203
204    if not out:
205        out.append(
206            {
207                "severity": "info",
208                "title": "nothing anomalous in the curves",
209                "evidence": "loss fell and eval did not diverge",
210                "suggestion": "scale up: the next question is the preset, not the schedule",
211            }
212        )
213    return out

Findings as {severity, title, evidence, suggestion} — suggestion names a config knob.

def render_markdown(runs: list[tuple[dict, list[dict]]]) -> str:
216def render_markdown(runs: list[tuple[dict, list[dict]]]) -> str:
217    rank = {"high": "🔴", "medium": "🟡", "info": "🔵"}
218    lines = ["# marola-sea training report", ""]
219    for s, findings in runs:
220        lines += [
221            f"## {s['name']} ({s['kind'].upper()})",
222            "",
223            f"- source: `{s['source']}`",
224            f"- steps: {s['steps']}, epochs run: {s['epochs_run'] or 'n/a'}",
225        ]
226        if s["train_loss"]:
227            lines.append(f"- train loss: {s['train_loss'][0]:.4f} → {s['train_loss'][-1]:.4f}")
228        if s["eval_loss"]:
229            lines.append(
230                f"- eval loss: {s['eval_loss'][0]:.4f} → {s['eval_loss'][-1]:.4f} "
231                f"(best {min(s['eval_loss']):.4f})"
232            )
233        if s["dpo_accuracy"]:
234            lines.append(f"- DPO reward accuracy: ends at {s['dpo_accuracy'][-1]:.2f}")
235        if s["runtime_s"]:
236            lines.append(
237                f"- runtime: {s['runtime_s'] / 60:.1f} min"
238                + (f", {s['samples_per_s']:.2f} samples/s" if s["samples_per_s"] else "")
239            )
240        lines += ["", "### What to change next run", ""]
241        for f in findings:
242            lines += [
243                f"**{rank.get(f['severity'], '•')} {f['title']}**  ",
244                f"{f['evidence']}  ",
245                f"→ {f['suggestion']}",
246                "",
247            ]
248    return "\n".join(lines)
def self_test() -> int:
251def self_test() -> int:
252    fails = 0
253
254    def ok(got, want, what):
255        nonlocal fails
256        if got == want:
257            print(f"  ok   {what}")
258        else:
259            print(f"  FAIL {what} — got {got!r} want {want!r}")
260            fails += 1
261
262    def titles(summary):
263        return [f["title"] for f in diagnose(summary)]
264
265    base = {
266        "name": "t",
267        "source": "x",
268        "kind": "sft",
269        "steps": 400,
270        "epochs_run": 3,
271        "train_loss": [],
272        "eval_loss": [],
273        "eval_epochs": [],
274        "grad_norm": [],
275        "dpo_accuracy": [],
276        "runtime_s": None,
277        "samples_per_s": None,
278    }
279
280    ok(
281        titles({**base, "steps": 12}),
282        ["the run was too short to learn much"],
283        "a 12-step run is flagged — the failure a small corpus actually hits",
284    )
285    ok(
286        "the run was too short to learn much" in titles({**base, "steps": 400}),
287        False,
288        "a 400-step run is not flagged for length",
289    )
290
291    over = {
292        **base,
293        "eval_loss": [2.0, 1.2, 1.5],
294        "eval_epochs": [1, 2, 3],
295        "train_loss": [3.0, 2.0, 1.0, 0.4],
296    }
297    ok(
298        "eval loss turned back up — overfitting" in titles(over),
299        True,
300        "eval loss rising while train loss falls is called overfitting",
301    )
302    ok(
303        "2" in diagnose(over)[0]["suggestion"],
304        True,
305        "and the suggestion names the epoch the run should have stopped at",
306    )
307
308    ok(
309        "train loss barely moved" in titles({**base, "train_loss": [2.0, 1.99, 1.98, 1.99]}),
310        True,
311        "a flat train curve is flagged as a learning-rate or dataset problem",
312    )
313    ok(
314        "train loss barely moved" in titles({**base, "train_loss": [3.0, 2.0, 1.0, 0.5]}),
315        False,
316        "a curve that fell by half is not",
317    )
318
319    spikes = {**base, "train_loss": [3.0, 2.0, 1.0, 0.5], "grad_norm": [0.5, 0.5, 0.6, 40.0, 0.5]}
320    ok("gradient-norm spikes" in titles(spikes), True, "a 40x gradient spike is flagged")
321    ok(
322        "gradient-norm spikes" in titles({**base, "grad_norm": [0.5, 0.6, 0.55, 0.7]}),
323        False,
324        "a steady gradient norm is not",
325    )
326
327    dpo = {
328        **base,
329        "kind": "dpo",
330        "train_loss": [0.7, 0.6, 0.5, 0.45],
331        "dpo_accuracy": [0.5, 0.52, 0.51],
332    }
333    ok(
334        "DPO is not separating chosen from rejected" in titles(dpo),
335        True,
336        "DPO reward accuracy at chance is flagged — the run learned no preference",
337    )
338    ok(
339        "DPO is not separating chosen from rejected"
340        in titles({**dpo, "dpo_accuracy": [0.8, 0.85, 0.9]}),
341        False,
342        "a DPO run that does separate them is not flagged",
343    )
344
345    ok(
346        titles({**base, "train_loss": [3.0, 2.0, 1.0, 0.5, 0.49, 0.49, 0.49]}),
347        ["nothing anomalous in the curves"],
348        "a healthy converged run reports nothing to change",
349    )
350
351    ok(
352        "still improving when it stopped — undertrained" in titles(over),
353        False,
354        "an overfitting run is never also told to train longer — the advice would contradict",
355    )
356    ok(
357        "still improving when it stopped — undertrained"
358        in titles({**base, "train_loss": [3.0, 2.5, 2.0, 1.5, 1.0]}),
359        True,
360        "but a run that is genuinely still falling, with no eval divergence, is",
361    )
362
363    ok(
364        len(parse_stdout("{'loss': 1.5, 'epoch': 0.5}\nnoise\n{'eval_loss': 1.2, 'epoch': 1.0}")),
365        2,
366        "stdout fallback recovers HF's single-quoted log dicts",
367    )
368    ok(
369        parse_stdout("{'loss': 1.5, 'epoch': 0.5}")[0]["loss"],
370        1.5,
371        "and their values survive the quote conversion",
372    )
373    ok(len(parse_stdout("no log lines here at all")), 0, "prose is not mistaken for a log record")
374
375    hist = [
376        {"loss": 2.0, "epoch": 0.5},
377        {"eval_loss": 1.8, "epoch": 1.0},
378        {"train_runtime": 120.0, "train_samples_per_second": 4.0},
379    ]
380    s = summarize(
381        {"name": "r", "source": "s", "log_history": hist, "epochs_run": 1, "global_step": 10}
382    )
383    ok(
384        (s["train_loss"], s["eval_loss"], s["runtime_s"]),
385        ([2.0], [1.8], 120.0),
386        "summarize splits train, eval and the final throughput record apart",
387    )
388    ok(s["kind"], "sft", "a run with no reward keys is classified as SFT")
389    ok(
390        summarize(
391            {
392                "name": "r",
393                "source": "s",
394                "log_history": [{DPO_ACC_KEY: 0.7, "epoch": 1}],
395                "epochs_run": 1,
396                "global_step": 1,
397            }
398        )["kind"],
399        "dpo",
400        "and one with them as DPO",
401    )
402
403    ok(
404        render_markdown([(s, diagnose(s))]).startswith("# marola-sea training report"),
405        True,
406        "the markdown report renders",
407    )
408
409    print(
410        "analyze_training self-test: ok"
411        if not fails
412        else f"analyze_training self-test: {fails} failure(s)",
413        file=sys.stderr if fails else sys.stdout,
414    )
415    return 1 if fails else 0
def main() -> int:
418def main() -> int:
419    ap = argparse.ArgumentParser(description=__doc__.split("\n")[0])
420    ap.add_argument("runs", nargs="*", type=Path, help="adapter dirs, or trainer_state.json files")
421    ap.add_argument("--markdown", type=Path, help="write the report here (default: stdout)")
422    ap.add_argument("--json", type=Path, help="also write the structured summary here")
423    ap.add_argument("--self-test", action="store_true")
424    args = ap.parse_args()
425
426    if args.self_test:
427        return self_test()
428    if not args.runs:
429        ap.error("give at least one run directory (or --self-test)")
430
431    analyzed = []
432    for target in args.runs:
433        try:
434            s = summarize(load_run(target))
435        except (FileNotFoundError, json.JSONDecodeError) as e:
436            print(f"analyze_training: skipping {target}: {e}", file=sys.stderr)
437            continue
438        analyzed.append((s, diagnose(s)))
439    if not analyzed:
440        print("analyze_training: no readable runs", file=sys.stderr)
441        return 1
442
443    report = render_markdown(analyzed)
444    if args.markdown:
445        args.markdown.write_text(report)
446        print(f"wrote {args.markdown}")
447    else:
448        print(report)
449    if args.json:
450        args.json.write_text(
451            json.dumps([{"summary": s, "findings": f} for s, f in analyzed], indent=2)
452        )
453        print(f"wrote {args.json}")
454    return 0