#!/usr/bin/env python3 """Launch and/or watch a Stage-1 training run: loss curve + benchmark-vs-base. This is the single "just run it" entrypoint for a training run you want to babysit. It reads the files run_sft.py writes (`/metrics.jsonl` and, when the in-training benchmark is enabled, `/eval_progress.jsonl`) and renders a compact snapshot: an ASCII loss curve, latest lr/grad-norm/eval-loss, and the held-out benchmark accuracy with its delta vs the base model. Modes: --once render one snapshot and exit (what an external watcher polls each tick) (default loop) re-render every --interval seconds (default 900 = 15 min) --launch CONFIG start `run_sft.py --config CONFIG` detached first, then watch The renderers are pure functions (stdlib only) so they are unit-tested offline; the training process itself is what needs the GPU. Examples: # launch Stage-1 and watch every 15 min, comparing to the base baseline: python training/scripts/train_watch.py \ --launch training/configs/stage1_lora_sft.yaml \ --output-dir /workspace/checkpoints/qwen36_27b_cybergym_stage1_lora_sft \ --base-eval reports/eval/base_eval.json --interval 900 # one snapshot of an already-running run: python training/scripts/train_watch.py \ --output-dir /workspace/checkpoints/qwen36_27b_cybergym_stage1_lora_sft --once """ from __future__ import annotations import argparse import json import subprocess import sys import time from pathlib import Path from typing import Any BLOCKS = "▁▂▃▄▅▆▇█" def read_jsonl(path: Path) -> list[dict[str, Any]]: if not path.is_file(): return [] rows = [] for line in path.read_text(encoding="utf-8").splitlines(): line = line.strip() if line: try: rows.append(json.loads(line)) except json.JSONDecodeError: pass return rows def sparkline(values: list[float], width: int = 60) -> str: vals = [v for v in values if isinstance(v, (int, float))] if not vals: return "(no data)" if len(vals) > width: # bucket-average down to width points step = len(vals) / width bucketed = [] for i in range(width): chunk = vals[int(i * step):int((i + 1) * step)] or [vals[min(int(i * step), len(vals) - 1)]] bucketed.append(sum(chunk) / len(chunk)) vals = bucketed lo, hi = min(vals), max(vals) if hi - lo < 1e-12: return BLOCKS[0] * len(vals) return "".join(BLOCKS[min(len(BLOCKS) - 1, int((v - lo) / (hi - lo) * (len(BLOCKS) - 1)))] for v in vals) def trend(values: list[float]) -> str: vals = [v for v in values if isinstance(v, (int, float))] if len(vals) < 2: return "n/a" delta = vals[-1] - vals[0] arrow = "↓" if delta < 0 else ("↑" if delta > 0 else "→") return f"{arrow} {delta:+.4f} over {len(vals)} pts" def fmt_age(ts: float | None) -> str: if not ts: return "n/a" secs = max(0, time.time() - ts) if secs < 90: return f"{int(secs)}s ago" if secs < 5400: return f"{secs/60:.1f}m ago" return f"{secs/3600:.1f}h ago" def summarize_metrics(rows: list[dict[str, Any]]) -> dict[str, Any]: train = [(r.get("step"), r.get("loss"), r.get("ts")) for r in rows if "loss" in r] evals = [(r.get("step"), r.get("eval_loss")) for r in rows if "eval_loss" in r] losses = [l for _, l, _ in train] last = rows[-1] if rows else {} return { "n_log": len(rows), "train_losses": losses, "first_loss": losses[0] if losses else None, "last_loss": losses[-1] if losses else None, "min_loss": min(losses) if losses else None, "eval_losses": [e for _, e in evals], "last_lr": last.get("learning_rate"), "last_grad_norm": last.get("grad_norm"), "last_epoch": last.get("epoch"), "last_step": last.get("step"), "last_ts": last.get("ts"), } def render_snapshot(output_dir: Path, metrics_rows, progress_rows, base_acc: dict[str, float]) -> str: m = summarize_metrics(metrics_rows) out: list[str] = [] epoch_str = f"{m['last_epoch']:.2f}" if isinstance(m["last_epoch"], (int, float)) else str(m["last_epoch"]) out.append(f"# Training watch — {output_dir.name}") out.append("") out.append(f"- log lines: {m['n_log']} | last step: {m['last_step']} epoch: {epoch_str}") out.append(f"- last update: {fmt_age(m['last_ts'])}") out.append("") out.append("## Train loss") out.append(f"`{sparkline(m['train_losses'])}`") out.append(f"- first {m['first_loss']} → last {m['last_loss']} (min {m['min_loss']})") out.append(f"- trend: {trend(m['train_losses'])}") if m["eval_losses"]: out.append("") out.append("## Eval loss") out.append(f"`{sparkline(m['eval_losses'])}` last {m['eval_losses'][-1]}") out.append("") out.append(f"- lr: {m['last_lr']} | grad_norm: {m['last_grad_norm']}") # checkpoints ckpts = sorted(p.name for p in output_dir.glob("checkpoint-*")) if output_dir.is_dir() else [] if ckpts: out.append(f"- checkpoints: {', '.join(ckpts)}") # per-set accuracy trajectory across checkpoints (early-trend detection) seqs = {} # set_name -> [(step, acc), ...] valid = [r for r in progress_rows if isinstance(r, dict) and "error" not in r] for r in valid: for name, met in (r.get("sets") or {}).items(): if isinstance(met, dict) and "accuracy" in met: seqs.setdefault(name, []).append((r.get("step"), met["accuracy"])) # overall early-trend verdict (loss + benchmark direction) loss_dir = trend(m["train_losses"]) # has ↓/↑ arrow bench_bits = [] for name, pts in seqs.items(): if len(pts) >= 2: d = pts[-1][1] - pts[0][1] arrow = "↑" if d > 0.01 else ("↓" if d < -0.01 else "→") bench_bits.append(f"{name.replace('_test','').replace('knowledge_','')}:{pts[0][1]:.0%}{arrow}{pts[-1][1]:.0%}") verdict = f"loss {loss_dir}" if bench_bits: verdict += " | " + " ".join(bench_bits) out.insert(3, f"- ⚡ EARLY TREND: {verdict}") out.append("") out.append("## Held-out benchmark vs base") if not valid: out.append("_(no in-training benchmark yet — runs every eval_every_steps + epoch end)_") else: latest = valid[-1] out.append(f"- as of step {latest.get('step')} (epoch {latest.get('epoch')}):") out.append("") out.append("| set | now | base | Δ vs base | trajectory |") out.append("|---|---|---|---|---|") for name, met in (latest.get("sets") or {}).items(): if not isinstance(met, dict) or "accuracy" not in met: continue kind = met.get("kind", "") b = base_acc.get(kind) acc = met["accuracy"] d = f"{acc - b:+.1%}" if isinstance(b, (int, float)) else "n/a" bs = f"{b:.0%}" if isinstance(b, (int, float)) else "n/a" traj = " ".join(f"{a:.0%}" for _, a in seqs.get(name, [])) f1 = f" F1 {met.get('f1_vuln',0):.2f}" if kind == "vuln_detection" else "" out.append(f"| {name} | {acc:.0%}{f1} | {bs} | {d} | {traj} |") out.append("") out.append(f"_rendered {fmt_age(time.time())[:-4] or 'now'} (utc epoch {int(time.time())})_") return "\n".join(out) def load_base_acc(base_eval: Path | None) -> dict[str, float]: if not base_eval or not base_eval.is_file(): return {} try: payload = json.loads(base_eval.read_text(encoding="utf-8")) return {r["kind"]: r["accuracy"] for r in payload.get("results", []) if "kind" in r} except Exception: return {} def one_snapshot(args) -> str: output_dir = Path(args.output_dir) metrics_path = Path(args.metrics) if args.metrics else output_dir / "metrics.jsonl" progress_path = Path(args.progress) if args.progress else output_dir / "eval_progress.jsonl" base_acc = load_base_acc(Path(args.base_eval) if args.base_eval else None) snap = render_snapshot(output_dir, read_jsonl(metrics_path), read_jsonl(progress_path), base_acc) if args.snapshot_out: out = Path(args.snapshot_out) else: out = Path("reports/training") / output_dir.name / "watch.md" out.parent.mkdir(parents=True, exist_ok=True) out.write_text(snap + "\n", encoding="utf-8") return snap def launch_training(config: str, log_path: Path) -> subprocess.Popen: log_path.parent.mkdir(parents=True, exist_ok=True) cmd = [sys.executable, str(Path(__file__).resolve().parent / "run_sft.py"), "--config", config] print(f"launching: {' '.join(cmd)} (log: {log_path})") return subprocess.Popen(cmd, stdout=log_path.open("w"), stderr=subprocess.STDOUT) def parse_args() -> argparse.Namespace: p = argparse.ArgumentParser(description=__doc__, formatter_class=argparse.RawDescriptionHelpFormatter) p.add_argument("--output-dir", required=True, help="Training output_dir (holds metrics.jsonl).") p.add_argument("--metrics", help="Override path to metrics.jsonl.") p.add_argument("--progress", help="Override path to eval_progress.jsonl.") p.add_argument("--base-eval", help="reports/eval/base_eval.json for vs-base comparison.") p.add_argument("--snapshot-out", help="Where to write the snapshot markdown.") p.add_argument("--interval", type=int, default=900, help="Seconds between snapshots (default 900).") p.add_argument("--once", action="store_true", help="Render one snapshot and exit.") p.add_argument("--max-ticks", type=int, default=0, help="Stop after N snapshots (0 = until training ends).") p.add_argument("--launch", help="Launch run_sft.py with this config before watching.") p.add_argument("--launch-log", default="/workspace/tmp/stage1_train.log") return p.parse_args() def main() -> int: args = parse_args() proc = None if args.launch: proc = launch_training(args.launch, Path(args.launch_log)) if args.once and not args.launch: print(one_snapshot(args)) return 0 tick = 0 while True: tick += 1 print("\n" + "=" * 72) print(one_snapshot(args)) if proc is not None and proc.poll() is not None: print(f"\n[training process exited rc={proc.returncode}] final snapshot above.") return proc.returncode or 0 if args.max_ticks and tick >= args.max_ticks: return 0 if args.once: return 0 time.sleep(args.interval) if __name__ == "__main__": raise SystemExit(main())