seed runs

This commit is contained in:
YannAhlgrim
2026-07-06 13:50:35 +02:00
parent d053e0a8e0
commit 56f15dd856
12 changed files with 716 additions and 57 deletions
+127 -57
View File
@@ -1,7 +1,9 @@
import os
import random
import shutil
import sys
import json
import time
import yaml
import logging
@@ -35,6 +37,22 @@ logging.basicConfig(stream=sys.stdout, level=logging.INFO)
logger = logging.getLogger()
def _set_seed(seed):
"""Seed all RNGs so multi-seed runs are reproducible and distinct."""
random.seed(seed)
np.random.seed(seed)
torch.manual_seed(seed)
if torch.cuda.is_available():
torch.cuda.manual_seed_all(seed)
def _format_hms(seconds):
seconds = int(round(seconds))
h, rem = divmod(seconds, 3600)
m, s = divmod(rem, 60)
return f"{h:02d}:{m:02d}:{s:02d}"
def strip_module_prefix(state_dict):
if not any(k.startswith("module.") for k in state_dict.keys()):
return state_dict
@@ -213,6 +231,10 @@ def main(args, resume_preempt=False):
l_args = args["logging"]
v_args = args["validation"]
es_args = o_args["early_stopping"]
seed = int(m_args.get("seed", _GLOBAL_SEED))
_set_seed(seed)
logger.info(f"Using seed {seed}")
accum_steps = o_args.get("gradient_accumulation_steps", 1)
use_gradient_checkpointing = m_args.get("use_gradient_checkpointing", True)
@@ -424,9 +446,13 @@ def main(args, resume_preempt=False):
if is_best:
torch.save(save_dict, best_path)
train_start_time = time.perf_counter()
epochs_run = 0
early_stopped = False
for epoch in range(start_epoch, o_args["epochs"]):
train_sampler.set_epoch(epoch)
val_sampler.set_epoch(epoch)
epochs_run = epoch + 1
model.train()
if o_args["freeze_weights"]:
@@ -501,12 +527,22 @@ def main(args, resume_preempt=False):
dist.broadcast(stop_tensor, src=0)
if bool(stop_tensor.item()):
early_stopped = True
if rank == 0:
logger.info(
f"Early stopping at epoch {epoch + 1}. Best val_loss={early_stopper.best_metric:.6f} @ epoch {early_stopper.best_epoch}"
)
break
train_time_seconds = time.perf_counter() - train_start_time
if rank == 0:
logger.info(
f"Training finished: epochs_run={epochs_run} "
f"early_stopped={early_stopped} "
f"train_time={_format_hms(train_time_seconds)} "
f"({train_time_seconds:.1f}s)"
)
if early_stopper.enabled and early_stopper.restore_best_weights:
if rank == 0:
logger.info("Restoring best model weights before exit")
@@ -523,73 +559,107 @@ def main(args, resume_preempt=False):
)
if rank == 0:
eval_args = {
"meta": {
"seed": m_args.get("seed", _GLOBAL_SEED),
"model_name": m_args["model_name"],
"embed_dim": m_args["embed_dim"],
"num_classes": m_args["num_classes"],
"patch_size": mk_args.get("patch_size", m_args.get("patch_size", 16)),
"crop_size": d_args.get("crop_size", m_args.get("crop_size", 224)),
"use_bfloat16": m_args.get("use_bfloat16", True),
"representation_type": representation_type,
"head_type": head_type,
"checkpoint_path": best_path,
"force_single_process": True,
},
"data": {
"batch_size": d_args.get("batch_size", 128),
"root_path": d_args.get("root_path", "./wilds_data"),
"num_workers": d_args.get("num_workers", 8),
"pin_mem": d_args.get("pin_mem", True),
"split": "test",
"download": True,
},
"logging": {
"write_tag": "iwildcam_test",
"auto_folder": True,
},
}
eval_result = eval_wilds_main(args=eval_args)
# Evaluate on both the in-distribution (id_test) and out-of-distribution
# (test) splits so the generalization gap can be measured.
# WILDS-iWildCam: "id_test" == Test ID, "test" == Test OOD.
eval_splits = [
("id_test", "iwildcam_id_test"),
("test", "iwildcam_test"),
]
def _make_eval_args(split, write_tag):
return {
"meta": {
"seed": seed,
"model_name": m_args["model_name"],
"embed_dim": m_args["embed_dim"],
"num_classes": m_args["num_classes"],
"patch_size": mk_args.get("patch_size", m_args.get("patch_size", 16)),
"crop_size": d_args.get("crop_size", m_args.get("crop_size", 224)),
"use_bfloat16": m_args.get("use_bfloat16", True),
"representation_type": representation_type,
"head_type": head_type,
"checkpoint_path": best_path,
"force_single_process": True,
},
"data": {
"batch_size": d_args.get("batch_size", 128),
"root_path": d_args.get("root_path", "./wilds_data"),
"num_workers": d_args.get("num_workers", 8),
"pin_mem": d_args.get("pin_mem", True),
"split": split,
"download": True,
},
"logging": {
"write_tag": write_tag,
"auto_folder": True,
},
}
eval_results = {}
eval_folder = None
for split, write_tag in eval_splits:
eval_args = _make_eval_args(split, write_tag)
result = eval_wilds_main(args=eval_args)
eval_results[split] = result
if eval_folder is None:
eval_folder = eval_args.get("logging", {}).get("folder")
# Common run-level info folded into every metrics JSON + the params summary.
run_info = {
"seed": seed,
"train_time_seconds": float(train_time_seconds),
"train_time_hms": _format_hms(train_time_seconds),
"epochs_run": int(epochs_run),
"configured_epochs": int(o_args["epochs"]),
"best_epoch": int(early_stopper.best_epoch),
"early_stopped": bool(early_stopped),
"best_val_loss": float(early_stopper.best_metric),
}
# Fold run_info into each split's metrics JSON so an aggregator can read
# metrics + seed + timing + epochs from a single file per split.
for split, result in eval_results.items():
if not result:
continue
metrics_path = result.get("metrics_path")
if not metrics_path or not os.path.exists(metrics_path):
continue
try:
with open(metrics_path, "r") as f:
metrics_obj = json.load(f)
metrics_obj["run_info"] = run_info
metrics_obj["split"] = split
with open(metrics_path, "w") as f:
json.dump(metrics_obj, f, indent=2, sort_keys=True)
except (OSError, json.JSONDecodeError):
logger.warning(f"Could not augment metrics JSON for split {split}")
eval_folder = eval_args.get("logging", {}).get("folder")
if eval_folder:
try:
params_out = yaml.safe_load(yaml.dump(args))
params_out.setdefault("meta", {})["representation_type"] = representation_type
params_out.setdefault("meta", {})["head_type"] = head_type
metric_key = m_args.get("selection_metric", "macro_f1")
eval_root = os.path.join("experiment_logs", "eval-wilds")
ranking = {
"metric_key": metric_key,
"this_run_name": os.path.basename(os.path.normpath(eval_folder)),
"this_is_best": None,
"this_metric": None,
"best_run_name": None,
"best_metric": None,
}
if os.path.isdir(eval_root):
rows = _collect_eval_rows(eval_root, metric_key)
if rows:
rows.sort(key=lambda r: r[0], reverse=True)
best_value, best_run_name, _ = rows[0]
this_run_name = ranking["this_run_name"]
this_value = None
for value, run_name, _ in rows:
if run_name == this_run_name:
this_value = value
break
ranking["this_metric"] = this_value
ranking["best_run_name"] = best_run_name
ranking["best_metric"] = best_value
if this_value is not None:
ranking["this_is_best"] = this_run_name == best_run_name
params_out["results"] = {
"best_val_loss": float(early_stopper.best_metric),
"best_epoch": int(early_stopper.best_epoch),
"best_checkpoint": best_path,
"eval_metrics": eval_result.get("metrics") if eval_result else None,
"ranking": ranking,
"seed": seed,
"train_time_seconds": float(train_time_seconds),
"train_time_hms": _format_hms(train_time_seconds),
"epochs_run": int(epochs_run),
"configured_epochs": int(o_args["epochs"]),
"early_stopped": bool(early_stopped),
"eval_metrics_id_test": (
eval_results.get("id_test", {}).get("metrics")
if eval_results.get("id_test")
else None
),
"eval_metrics_test": (
eval_results.get("test", {}).get("metrics")
if eval_results.get("test")
else None
),
}
with open(os.path.join(eval_folder, "params-supervised.yaml"), "w") as f:
yaml.dump(params_out, f)