#!/usr/bin/env python3 """ BabyLM Challenge 2026 - Results Collector & Comparator Scans experiment checkpoints and evaluation results, produces a summary table. Usage: python collect_results.py # Show all results python collect_results.py --phase 1 2 # Filter by phase python collect_results.py --sort blimp # Sort by metric python collect_results.py --csv results.csv # Export to CSV """ import argparse import csv import json import re import sys from pathlib import Path ROOT = Path(__file__).resolve().parent CHECKPOINTS_DIR = ROOT / "checkpoints" RESULTS_DIR = ROOT / "experiment_results" RESULTS_JSON = RESULTS_DIR / "all_results.json" EXPERIMENTS_CSV = ROOT / "experiments.csv" def load_experiment_info() -> dict: """Load experiment metadata from CSV.""" info = {} with open(EXPERIMENTS_CSV) as f: reader = csv.DictReader(f) for row in reader: if row.get("id", "").startswith("##"): continue info[row["id"]] = row return info def load_run_results() -> dict: """Load run results from JSON (status, timing, etc.).""" if RESULTS_JSON.exists(): with open(RESULTS_JSON) as f: return json.load(f) return {} def find_checkpoints() -> dict: """Find all experiment checkpoints and extract training loss.""" results = {} if not CHECKPOINTS_DIR.exists(): return results for exp_dir in sorted(CHECKPOINTS_DIR.iterdir()): if not exp_dir.is_dir(): continue exp_name = exp_dir.name # Find latest checkpoint ckpts = sorted(exp_dir.glob("checkpoint_epoch*.pt")) if not ckpts: continue # Extract info from latest checkpoint filename latest = ckpts[-1] epoch_match = re.search(r"epoch(\d+)", latest.name) epoch = int(epoch_match.group(1)) if epoch_match else 0 results[exp_name] = { "checkpoint_dir": str(exp_dir), "latest_checkpoint": str(latest), "num_checkpoints": len(ckpts), "latest_epoch": epoch, } return results def parse_training_log(log_path: Path) -> dict: """Extract final training loss and other metrics from a log file.""" metrics = {} if not log_path.exists(): return metrics text = log_path.read_text() # Find last loss value: "Loss X.XXXX" loss_matches = re.findall(r"Loss (\d+\.\d+)", text) if loss_matches: metrics["final_loss"] = float(loss_matches[-1]) # Find parameter count param_match = re.search(r"Parameters: ([\d,]+)", text) if param_match: metrics["params"] = param_match.group(1) # Find tokens/sec tok_matches = re.findall(r"(\d+) tok/s", text) if tok_matches: metrics["tok_per_sec"] = int(tok_matches[-1]) # Find training time per epoch time_matches = re.findall(r"Time (\d+)s", text) if time_matches: metrics["total_time_sec"] = sum(int(t) for t in time_matches) return metrics def collect_all() -> list[dict]: """Collect all results into a unified list.""" exp_info = load_experiment_info() run_results = load_run_results() checkpoints = find_checkpoints() rows = [] for exp_id, info in exp_info.items(): row = { "id": exp_id, "phase": info.get("phase", ""), "name": info.get("name", ""), "arch": info.get("arch", ""), "epochs": info.get("epochs", ""), "data": info.get("data", ""), } # Run status run = run_results.get(exp_id, {}) row["status"] = run.get("status", "not_started") row["elapsed_sec"] = run.get("elapsed_sec", "") # Training log metrics log_path = RESULTS_DIR / f"{exp_id.replace('.', '_')}.log" log_metrics = parse_training_log(log_path) row["final_loss"] = log_metrics.get("final_loss", "") row["params"] = log_metrics.get("params", "") row["tok_per_sec"] = log_metrics.get("tok_per_sec", "") # Checkpoint info ckpt = checkpoints.get(info.get("name", ""), {}) row["checkpoints"] = ckpt.get("num_checkpoints", 0) # TODO: evaluation metrics (BLiMP, EWoK, GLUE, etc.) # These will be populated after running evaluation pipeline row["blimp"] = "" row["ewok"] = "" row["glue"] = "" row["entity_tracking"] = "" rows.append(row) return rows def print_table(rows: list[dict], sort_by: str = None): """Print a formatted results table.""" if sort_by and rows: rows = sorted(rows, key=lambda r: (r.get(sort_by) or 0), reverse=True) # Header fmt = "{:<10} {:<6} {:<28} {:<10} {:<8} {:<10} {:<8} {:<6} {:<6} {:<6}" header = fmt.format("ID", "Phase", "Name", "Status", "Loss", "Params", "BLiMP", "EWoK", "GLUE", "ET") print(header) print("-" * len(header)) for row in rows: loss = f"{row['final_loss']:.4f}" if row.get("final_loss") else "" print(fmt.format( row["id"], row["phase"], row["name"][:28], row["status"][:10], loss, str(row.get("params", ""))[:10], str(row.get("blimp", "")), str(row.get("ewok", "")), str(row.get("glue", "")), str(row.get("entity_tracking", "")), )) def export_csv(rows: list[dict], output_path: str): """Export results to CSV.""" if not rows: return keys = rows[0].keys() with open(output_path, "w", newline="") as f: writer = csv.DictWriter(f, fieldnames=keys) writer.writeheader() writer.writerows(rows) print(f"\nExported {len(rows)} results to {output_path}") def print_summary(rows: list[dict]): """Print summary statistics.""" total = len(rows) by_status = {} for r in rows: s = r.get("status", "unknown") by_status[s] = by_status.get(s, 0) + 1 print(f"\n{'='*40}") print(f" Total experiments: {total}") for status, count in sorted(by_status.items()): pct = count / total * 100 print(f" {status:<15} {count:>4} ({pct:.0f}%)") print(f"{'='*40}") def main(): parser = argparse.ArgumentParser(description="BabyLM Results Collector") parser.add_argument("--phase", nargs="*", help="Filter by phase") parser.add_argument("--sort", default=None, help="Sort by column (e.g., blimp, final_loss)") parser.add_argument("--csv", default=None, help="Export to CSV file") parser.add_argument("--completed", action="store_true", help="Show only completed experiments") args = parser.parse_args() rows = collect_all() # Filter if args.phase: rows = [r for r in rows if r["phase"] in args.phase] if args.completed: rows = [r for r in rows if r["status"] == "completed"] # Display print_table(rows, sort_by=args.sort) print_summary(rows) # Export if args.csv: export_csv(rows, args.csv) if __name__ == "__main__": main()