From 8956531fe5df1f6fceeedb40b2243cc472ab395a Mon Sep 17 00:00:00 2001 From: YannAhlgrim Date: Sun, 5 Apr 2026 18:44:54 +0200 Subject: [PATCH] change explicitly for slurm cluster --- main_distributed.py | 70 +++++++++++++++++++++++---------------------- 1 file changed, 36 insertions(+), 34 deletions(-) diff --git a/main_distributed.py b/main_distributed.py index 7cb3846..9941c73 100644 --- a/main_distributed.py +++ b/main_distributed.py @@ -21,46 +21,43 @@ logger = logging.getLogger() parser = argparse.ArgumentParser() +parser.add_argument("--folder", type=str, help="location to save submitit logs") parser.add_argument( - '--folder', type=str, - help='location to save submitit logs') + "--batch-launch", + action="store_true", + help="whether fname points to a file to batch-lauch several config files", +) parser.add_argument( - '--batch-launch', action='store_true', - help='whether fname points to a file to batch-lauch several config files') + "--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( - '--fname', type=str, - help='yaml file containing config file names to launch', - default='configs.yaml') + "--nodes", type=int, default=1, help="num. nodes to request for job" +) 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') + "--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): + 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}') + logger.info(f"called-params {fname}") # -- load script params params = None - with open(fname, 'r') as y_file: + with open(fname, "r") as y_file: params = yaml.load(y_file, Loader=yaml.FullLoader) - logger.info('loaded params...') + logger.info("loaded params...") pp = pprint.PrettyPrinter(indent=4) pp.pprint(params) @@ -69,21 +66,24 @@ class Trainer: def checkpoint(self): fb_trainer = Trainer(self.fname, True) - return submitit.helpers.DelayedSubmission(fb_trainer,) + return submitit.helpers.DelayedSubmission( + fb_trainer, + ) def launch(): - executor = submitit.AutoExecutor( - folder=os.path.join(args.folder, 'job_%j'), - slurm_max_num_timeout=20) + executor = submitit.SlurmExecutor( + folder=os.path.join(args.folder, "job_%j"), max_num_timeout=20 + ) executor.update_parameters( - slurm_partition=args.partition, - slurm_mem_per_gpu='55G', - timeout_min=args.time, + partition=args.partition, + mem_per_gpu="55G", + time=args.time, nodes=args.nodes, - tasks_per_node=args.tasks_per_node, + ntasks_per_node=args.tasks_per_node, cpus_per_task=10, - gpus_per_node=args.tasks_per_node) + gpus_per_node=args.tasks_per_node, + ) config_fnames = [args.fname] @@ -91,7 +91,9 @@ def launch(): with executor.batch(): for cf in config_fnames: fb_trainer = Trainer(cf) - job = executor.submit(fb_trainer,) + job = executor.submit( + fb_trainer, + ) trainers.append(fb_trainer) jobs.append(job) @@ -99,6 +101,6 @@ def launch(): print(job.job_id) -if __name__ == '__main__': +if __name__ == "__main__": args = parser.parse_args() launch()