add grid search for supervised pipeline
This commit is contained in:
@@ -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