add grid search for supervised pipeline

This commit is contained in:
YannAhlgrim
2026-05-23 13:40:46 +02:00
parent 5240c4f3ff
commit 34a15d646d
5 changed files with 235 additions and 10 deletions
+16
View File
@@ -0,0 +1,16 @@
base_config: configs/supervised_wilds_vith14_ep300-lp.yaml
constants:
logging.write_tag: linear_probe
grid:
optimization.lr: [0.01, 0.05]
optimization.weight_decay: [0.0, 5.0e-4]
data.batch_size: [128, 256]
launch:
folder: submitit_logs/
partition: gpu1
nodes: 1
tasks_per_node: 1
time: 300
+8
View File
@@ -497,6 +497,14 @@ def main(args, resume_preempt=False):
} }
eval_wilds_main(args=eval_args) eval_wilds_main(args=eval_args)
if os.path.exists(folder):
for fname in os.listdir(folder):
if fname.endswith(".pth.tar"):
try:
os.remove(os.path.join(folder, fname))
except OSError:
logger.warning(f"Could not remove checkpoint file: {fname}")
if __name__ == "__main__": if __name__ == "__main__":
raise RuntimeError( raise RuntimeError(
+31 -10
View File
@@ -34,6 +34,10 @@ def _format_value(value):
return None return None
if isinstance(value, float): if isinstance(value, float):
return f"{value:g}" return f"{value:g}"
if isinstance(value, (list, tuple)):
return "-".join(_format_value(v) for v in value)
if isinstance(value, bool):
return "1" if value else "0"
return str(value) return str(value)
@@ -42,6 +46,7 @@ def build_run_name(args):
data_args = args.get("data", {}) data_args = args.get("data", {})
opt_args = args.get("optimization", {}) opt_args = args.get("optimization", {})
mask_args = args.get("mask", {}) mask_args = args.get("mask", {})
val_args = args.get("validation", {})
model_name = meta_args.get("model_name", "model") model_name = meta_args.get("model_name", "model")
patch_size = mask_args.get("patch_size", meta_args.get("patch_size")) patch_size = mask_args.get("patch_size", meta_args.get("patch_size"))
@@ -52,16 +57,32 @@ def build_run_name(args):
weight_decay = opt_args.get("weight_decay") weight_decay = opt_args.get("weight_decay")
epochs = opt_args.get("epochs") epochs = opt_args.get("epochs")
parts = [ parts = [model_name]
model_name,
f"p{_format_value(patch_size)}" if patch_size is not None else None, def add(prefix, value):
f"c{_format_value(crop_size)}" if crop_size is not None else None, formatted = _format_value(value)
f"bs{_format_value(batch_size)}" if batch_size is not None else None, if formatted is not None:
str(optimizer).lower(), parts.append(f"{prefix}{formatted}")
f"lr{_format_value(lr)}" if lr is not None else None,
f"wd{_format_value(weight_decay)}" if weight_decay is not None else None, add("p", patch_size)
f"ep{_format_value(epochs)}" if epochs is not None else None, add("c", crop_size)
] add("bs", batch_size)
parts.append(str(optimizer).lower())
add("lr", lr)
add("wd", weight_decay)
add("ep", epochs)
add("sched", opt_args.get("lr_schedule"))
add("cos", opt_args.get("use_cosine_schedule"))
add("ms", opt_args.get("step_milestones"))
add("sg", opt_args.get("step_gamma"))
add("wu", opt_args.get("warmup"))
add("mom", opt_args.get("momentum"))
add("leta", opt_args.get("lars_eta"))
add("leps", opt_args.get("lars_eps"))
add("ipe", opt_args.get("ipe_scale"))
add("cs", data_args.get("crop_scale"))
add("eval", val_args.get("eval_every"))
return "-".join([p for p in parts if p]) return "-".join([p for p in parts if p])
+110
View File
@@ -0,0 +1,110 @@
import argparse
import itertools
import os
import tempfile
import yaml
import submitit
from src.train_supervised import main as app_main
def _set_by_dotted_key(config, dotted_key, value):
keys = dotted_key.split(".")
cur = config
for key in keys[:-1]:
if key not in cur or not isinstance(cur[key], dict):
cur[key] = {}
cur = cur[key]
cur[keys[-1]] = value
def _deep_update(base, updates):
for key, value in updates.items():
_set_by_dotted_key(base, key, value)
return base
def _load_yaml(path):
with open(path, "r") as f:
return yaml.load(f, Loader=yaml.FullLoader)
def _expand_grid(grid_dict):
keys = list(grid_dict.keys())
values = [grid_dict[k] for k in keys]
for combo in itertools.product(*values):
yield dict(zip(keys, combo))
class GridTrainer:
def __init__(self, fname):
self.fname = fname
def __call__(self):
with open(self.fname, "r") as y_file:
params = yaml.load(y_file, Loader=yaml.FullLoader)
app_main(args=params, resume_preempt=False)
def main():
parser = argparse.ArgumentParser()
parser.add_argument("--grid", required=True, help="grid yaml file")
parser.add_argument("--folder", type=str, help="location to save submitit logs")
parser.add_argument("--partition", type=str, help="cluster partition to submit jobs on")
parser.add_argument("--nodes", type=int, default=1, help="num. nodes to request for job")
parser.add_argument(
"--tasks-per-node", type=int, default=1, help="num. procs to per node"
)
parser.add_argument("--time", type=int, default=4300, help="time in minutes to run job")
args = parser.parse_args()
grid_cfg = _load_yaml(args.grid)
base_config_path = grid_cfg["base_config"]
base_params = _load_yaml(base_config_path)
constants = grid_cfg.get("constants", {})
grid = grid_cfg.get("grid", {})
launch = grid_cfg.get("launch", {})
params_list = []
for overrides in _expand_grid(grid):
params = yaml.safe_load(yaml.dump(base_params))
_deep_update(params, constants)
_deep_update(params, overrides)
params_list.append(params)
log_folder = args.folder or launch.get("folder")
if not log_folder:
raise ValueError("submitit folder required via --folder or grid launch.folder")
executor = submitit.SlurmExecutor(
folder=os.path.join(log_folder, "job_%j"), max_num_timeout=20
)
executor.update_parameters(
partition=args.partition or launch.get("partition"),
mem_per_gpu=launch.get("mem_per_gpu", "55G"),
time=args.time or int(launch.get("time", 4300)),
nodes=args.nodes or int(launch.get("nodes", 1)),
ntasks_per_node=args.tasks_per_node or int(launch.get("tasks_per_node", 1)),
cpus_per_task=int(launch.get("cpus_per_task", 10)),
gpus_per_node=args.tasks_per_node or int(launch.get("tasks_per_node", 1)),
)
temp_dir = tempfile.mkdtemp(prefix="grid_configs_")
jobs = []
with executor.batch():
for idx, params in enumerate(params_list):
tmp_path = os.path.join(temp_dir, f"grid_{idx}.yaml")
with open(tmp_path, "w") as f:
yaml.dump(params, f)
job = executor.submit(GridTrainer(tmp_path))
jobs.append(job)
for job in jobs:
print(job.job_id)
if __name__ == "__main__":
main()
+70
View File
@@ -0,0 +1,70 @@
import argparse
import json
import os
def _find_metric(metrics, key):
if isinstance(metrics, dict):
if key in metrics:
return metrics[key]
for value in metrics.values():
found = _find_metric(value, key)
if found is not None:
return found
return None
def _collect_metrics(root_dir, key):
rows = []
for dirpath, _, filenames in os.walk(root_dir):
for fname in filenames:
if not fname.endswith("_metrics.json"):
continue
path = os.path.join(dirpath, fname)
try:
with open(path, "r") as f:
metrics = json.load(f)
except (OSError, json.JSONDecodeError):
continue
value = _find_metric(metrics, key)
if value is None:
continue
run_name = os.path.basename(os.path.dirname(path))
rows.append((float(value), run_name, path))
return rows
def main():
parser = argparse.ArgumentParser()
parser.add_argument(
"--root",
default="experiment_logs/eval-wilds",
help="root folder for eval logs",
)
parser.add_argument(
"--metric",
default="macro_f1",
help="metric key to rank by",
)
parser.add_argument(
"--top",
type=int,
default=20,
help="number of runs to display",
)
args = parser.parse_args()
rows = _collect_metrics(args.root, args.metric)
rows.sort(key=lambda r: r[0], reverse=True)
if not rows:
print("No metrics found.")
return
print(f"Ranking by '{args.metric}' (top {min(args.top, len(rows))})")
for idx, (value, run_name, path) in enumerate(rows[: args.top], start=1):
print(f"{idx:>3} | {value:.6f} | {run_name} | {path}")
if __name__ == "__main__":
main()