add supervised training pipeline
This commit is contained in:
@@ -0,0 +1,106 @@
|
||||
# Copyright (c) Meta Platforms, Inc. and affiliates.
|
||||
# All rights reserved.
|
||||
#
|
||||
# This source code is licensed under the license found in the
|
||||
# LICENSE file in the root directory of this source tree.
|
||||
#
|
||||
|
||||
import argparse
|
||||
import logging
|
||||
import os
|
||||
import pprint
|
||||
import sys
|
||||
import yaml
|
||||
|
||||
import submitit
|
||||
|
||||
from src.train_supervised import main as app_main
|
||||
|
||||
logging.basicConfig(stream=sys.stdout, level=logging.INFO)
|
||||
logger = logging.getLogger()
|
||||
|
||||
|
||||
parser = argparse.ArgumentParser()
|
||||
parser.add_argument("--folder", type=str, help="location to save submitit logs")
|
||||
parser.add_argument(
|
||||
"--batch-launch",
|
||||
action="store_true",
|
||||
help="whether fname points to a file to batch-lauch several config files",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--fname",
|
||||
type=str,
|
||||
help="yaml file containing config file names to launch",
|
||||
default="configs.yaml",
|
||||
)
|
||||
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")
|
||||
|
||||
|
||||
class Trainer:
|
||||
def __init__(self, fname="configs.yaml", load_model=None):
|
||||
self.fname = fname
|
||||
self.load_model = load_model
|
||||
|
||||
def __call__(self):
|
||||
fname = self.fname
|
||||
load_model = self.load_model
|
||||
logger.info(f"called-params {fname}")
|
||||
|
||||
# -- load script params
|
||||
params = None
|
||||
with open(fname, "r") as y_file:
|
||||
params = yaml.load(y_file, Loader=yaml.FullLoader)
|
||||
logger.info("loaded params...")
|
||||
pp = pprint.PrettyPrinter(indent=4)
|
||||
pp.pprint(params)
|
||||
|
||||
resume_preempt = False if load_model is None else load_model
|
||||
app_main(args=params, resume_preempt=resume_preempt)
|
||||
|
||||
def checkpoint(self):
|
||||
fb_trainer = Trainer(self.fname, True)
|
||||
return submitit.helpers.DelayedSubmission(
|
||||
fb_trainer,
|
||||
)
|
||||
|
||||
|
||||
def launch():
|
||||
executor = submitit.SlurmExecutor(
|
||||
folder=os.path.join(args.folder, "job_%j"), max_num_timeout=20
|
||||
)
|
||||
executor.update_parameters(
|
||||
partition=args.partition,
|
||||
mem_per_gpu="55G",
|
||||
time=args.time,
|
||||
nodes=args.nodes,
|
||||
ntasks_per_node=args.tasks_per_node,
|
||||
cpus_per_task=10,
|
||||
gpus_per_node=args.tasks_per_node,
|
||||
)
|
||||
|
||||
config_fnames = [args.fname]
|
||||
|
||||
jobs, trainers = [], []
|
||||
with executor.batch():
|
||||
for cf in config_fnames:
|
||||
fb_trainer = Trainer(cf)
|
||||
job = executor.submit(
|
||||
fb_trainer,
|
||||
)
|
||||
trainers.append(fb_trainer)
|
||||
jobs.append(job)
|
||||
|
||||
for job in jobs:
|
||||
print(job.job_id)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
args = parser.parse_args()
|
||||
launch()
|
||||
Reference in New Issue
Block a user