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)
|
||||
|
||||
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(
|
||||
|
||||
+31
-10
@@ -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])
|
||||
|
||||
|
||||
|
||||
@@ -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