add grid search for supervised pipeline
This commit is contained in:
@@ -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
|
||||||
@@ -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
@@ -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])
|
||||||
|
|
||||||
|
|
||||||
|
|||||||
@@ -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()
|
||||||
@@ -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()
|
||||||
Reference in New Issue
Block a user