From 34a15d646d3e53ca7a0e7a873f4e7044177e971b Mon Sep 17 00:00:00 2001 From: YannAhlgrim Date: Sat, 23 May 2026 13:40:46 +0200 Subject: [PATCH] add grid search for supervised pipeline --- configs/grids/lp_grid.yaml | 16 ++++++ src/train_supervised.py | 8 +++ src/utils/logging.py | 41 ++++++++++---- tools/run_grid.py | 110 +++++++++++++++++++++++++++++++++++++ tools/summarize_grid.py | 70 +++++++++++++++++++++++ 5 files changed, 235 insertions(+), 10 deletions(-) create mode 100644 configs/grids/lp_grid.yaml create mode 100644 tools/run_grid.py create mode 100644 tools/summarize_grid.py diff --git a/configs/grids/lp_grid.yaml b/configs/grids/lp_grid.yaml new file mode 100644 index 0000000..c5a32ec --- /dev/null +++ b/configs/grids/lp_grid.yaml @@ -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 diff --git a/src/train_supervised.py b/src/train_supervised.py index 42a4400..f92529e 100644 --- a/src/train_supervised.py +++ b/src/train_supervised.py @@ -497,6 +497,14 @@ def main(args, resume_preempt=False): } 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__": raise RuntimeError( diff --git a/src/utils/logging.py b/src/utils/logging.py index 4a2e794..b82c198 100644 --- a/src/utils/logging.py +++ b/src/utils/logging.py @@ -34,6 +34,10 @@ def _format_value(value): return None if isinstance(value, float): 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) @@ -42,6 +46,7 @@ def build_run_name(args): data_args = args.get("data", {}) opt_args = args.get("optimization", {}) mask_args = args.get("mask", {}) + val_args = args.get("validation", {}) model_name = meta_args.get("model_name", "model") 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") epochs = opt_args.get("epochs") - parts = [ - model_name, - f"p{_format_value(patch_size)}" if patch_size is not None else None, - f"c{_format_value(crop_size)}" if crop_size is not None else None, - f"bs{_format_value(batch_size)}" if batch_size is not None else None, - str(optimizer).lower(), - 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, - f"ep{_format_value(epochs)}" if epochs is not None else None, - ] + parts = [model_name] + + def add(prefix, value): + formatted = _format_value(value) + if formatted is not None: + parts.append(f"{prefix}{formatted}") + + add("p", patch_size) + 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]) diff --git a/tools/run_grid.py b/tools/run_grid.py new file mode 100644 index 0000000..b3ffb88 --- /dev/null +++ b/tools/run_grid.py @@ -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() diff --git a/tools/summarize_grid.py b/tools/summarize_grid.py new file mode 100644 index 0000000..2f5b677 --- /dev/null +++ b/tools/summarize_grid.py @@ -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()