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())
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.
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.
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}")
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.
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 }
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.
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)
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
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