min seeds

This commit is contained in:
YannAhlgrim
2026-07-07 09:43:28 +02:00
parent 20bf9515f2
commit 1a3916ce1a
+102 -7
View File
@@ -22,6 +22,7 @@ Outputs:
Usage: Usage:
python3 tools/aggregate_seeds.py --root experiment_logs/eval-wilds python3 tools/aggregate_seeds.py --root experiment_logs/eval-wilds
python3 tools/aggregate_seeds.py --primary F1-macro_all --acc acc_avg python3 tools/aggregate_seeds.py --primary F1-macro_all --acc acc_avg
python3 tools/aggregate_seeds.py --root experiment_logs/eval-wilds --min-seeds 5
""" """
import argparse import argparse
import csv import csv
@@ -60,6 +61,67 @@ def _strip_seed(run_name):
return _SEED_SUFFIX_RE.sub("", run_name) return _SEED_SUFFIX_RE.sub("", run_name)
def _extract_run_info(*metrics_objs):
"""Find a run_info dict inside any metrics object (dict- or list-form).
Eval-only runs write metrics as a list ([{metrics}, "summary"]) with no
run_info; training runs write a dict with a top-level run_info. Search both.
"""
def _search(obj):
if isinstance(obj, dict):
ri = obj.get("run_info")
if isinstance(ri, dict):
return ri
for v in obj.values():
found = _search(v)
if found is not None:
return found
elif isinstance(obj, list):
for item in obj:
found = _search(item)
if found is not None:
return found
return None
for m in metrics_objs:
found = _search(m)
if found is not None:
return found
return {}
def _seed_from_name(run_name):
"""Parse a trailing -seedN from a run folder name; None if absent."""
m = re.search(r"-seed(\d+)$", run_name)
return int(m.group(1)) if m else None
def _seed_from_params(run_dir):
"""Read `seed:` from params-eval.yaml without requiring PyYAML."""
path = os.path.join(run_dir, "params-eval.yaml")
try:
with open(path, "r") as f:
for line in f:
m = re.match(r"\s*seed\s*:\s*(\d+)\s*$", line)
if m:
return int(m.group(1))
except OSError:
pass
return None
def _resolve_seed(run_info, run_name, run_dir):
"""Seed detection chain: run_info -> folder name -> params-eval.yaml."""
seed = run_info.get("seed")
if seed is not None:
return seed
seed = _seed_from_name(run_name)
if seed is not None:
return seed
return _seed_from_params(run_dir)
def _load_json(path): def _load_json(path):
try: try:
with open(path, "r") as f: with open(path, "r") as f:
@@ -134,6 +196,12 @@ def main():
default="acc_avg", default="acc_avg",
help="average-accuracy metric key (default: %(default)s)", help="average-accuracy metric key (default: %(default)s)",
) )
parser.add_argument(
"--min-seeds",
type=int,
default=1,
help="only report models with at least this many seeds (default: %(default)s)",
)
args = parser.parse_args() args = parser.parse_args()
if not os.path.isdir(args.root): if not os.path.isdir(args.root):
@@ -151,17 +219,13 @@ def main():
if id_metrics is None and ood_metrics is None: if id_metrics is None and ood_metrics is None:
continue continue
run_info = {} run_info = _extract_run_info(ood_metrics, id_metrics)
for m in (ood_metrics, id_metrics):
if isinstance(m, dict) and isinstance(m.get("run_info"), dict):
run_info = m["run_info"]
break
group = _strip_seed(entry) group = _strip_seed(entry)
groups.setdefault(group, []).append( groups.setdefault(group, []).append(
{ {
"run_name": entry, "run_name": entry,
"seed": run_info.get("seed"), "seed": _resolve_seed(run_info, entry, run_dir),
"id_metrics": id_metrics, "id_metrics": id_metrics,
"ood_metrics": ood_metrics, "ood_metrics": ood_metrics,
"run_info": run_info, "run_info": run_info,
@@ -172,6 +236,26 @@ def main():
print(f"No metrics found under {args.root}") print(f"No metrics found under {args.root}")
return return
if args.min_seeds > 1:
dropped = {
name: len(recs)
for name, recs in groups.items()
if len(recs) < args.min_seeds
}
groups = {
name: recs for name, recs in groups.items() if len(recs) >= args.min_seeds
}
if dropped:
print(
f"Skipping {len(dropped)} model(s) with fewer than "
f"{args.min_seeds} seed(s):"
)
for name in sorted(dropped):
print(f" {name} ({dropped[name]} seed(s))")
if not groups:
print(f"No models have >= {args.min_seeds} seeds under {args.root}")
return
os.makedirs(args.out, exist_ok=True) os.makedirs(args.out, exist_ok=True)
csv_rows = [] csv_rows = []
@@ -259,10 +343,21 @@ def main():
gpu_reserved_vals = [r["run_info"].get("peak_gpu_reserved_gb") for r in records] gpu_reserved_vals = [r["run_info"].get("peak_gpu_reserved_gb") for r in records]
time_mean, time_std, _ = _mean_std(time_vals) time_mean, time_std, _ = _mean_std(time_vals)
epoch_mean, epoch_std, _ = _mean_std(epoch_vals) epoch_mean, epoch_std, _ = _mean_std(epoch_vals)
host_ram_mean, host_ram_std, _ = _mean_std(host_ram_vals) host_ram_mean, host_ram_std, host_ram_n = _mean_std(host_ram_vals)
gpu_alloc_mean, gpu_alloc_std, _ = _mean_std(gpu_alloc_vals) gpu_alloc_mean, gpu_alloc_std, _ = _mean_std(gpu_alloc_vals)
gpu_reserved_mean, gpu_reserved_std, _ = _mean_std(gpu_reserved_vals) gpu_reserved_mean, gpu_reserved_std, _ = _mean_std(gpu_reserved_vals)
if host_ram_n == 0:
# No usable peak_host_ram_gb in any seed's run_info. This usually means
# the metrics JSONs were (re)generated by a standalone eval run that did
# not fold in train_supervised.py's run_info, or run_info is absent.
missing_run_info = sum(1 for r in records if not r["run_info"])
print(
f"[warn] {group_name}: RAM was not recorded for these eval runs "
f"(no peak_host_ram_gb across {len(records)} seed(s); "
f"{missing_run_info} missing run_info entirely). RAM column stays blank."
)
# Headline (leaderboard) numbers. # Headline (leaderboard) numbers.
id_f1_mean, id_f1_std, _ = _mean_std(metric_values("id_metrics", args.primary)) id_f1_mean, id_f1_std, _ = _mean_std(metric_values("id_metrics", args.primary))
id_acc_mean, id_acc_std, _ = _mean_std(metric_values("id_metrics", args.acc)) id_acc_mean, id_acc_std, _ = _mean_std(metric_values("id_metrics", args.acc))