change explicitly for slurm cluster

This commit is contained in:
YannAhlgrim
2026-04-05 18:44:54 +02:00
parent 5da7b0f521
commit 8956531fe5
+36 -34
View File
@@ -21,46 +21,43 @@ logger = logging.getLogger()
parser = argparse.ArgumentParser() parser = argparse.ArgumentParser()
parser.add_argument("--folder", type=str, help="location to save submitit logs")
parser.add_argument( parser.add_argument(
'--folder', type=str, "--batch-launch",
help='location to save submitit logs') action="store_true",
help="whether fname points to a file to batch-lauch several config files",
)
parser.add_argument( parser.add_argument(
'--batch-launch', action='store_true', "--fname",
help='whether fname points to a file to batch-lauch several config files') 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( parser.add_argument(
'--fname', type=str, "--nodes", type=int, default=1, help="num. nodes to request for job"
help='yaml file containing config file names to launch', )
default='configs.yaml')
parser.add_argument( parser.add_argument(
'--partition', type=str, "--tasks-per-node", type=int, default=1, help="num. procs to per node"
help='cluster partition to submit jobs on') )
parser.add_argument( parser.add_argument("--time", type=int, default=4300, help="time in minutes to run job")
'--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: class Trainer:
def __init__(self, fname="configs.yaml", load_model=None):
def __init__(self, fname='configs.yaml', load_model=None):
self.fname = fname self.fname = fname
self.load_model = load_model self.load_model = load_model
def __call__(self): def __call__(self):
fname = self.fname fname = self.fname
load_model = self.load_model load_model = self.load_model
logger.info(f'called-params {fname}') logger.info(f"called-params {fname}")
# -- load script params # -- load script params
params = None params = None
with open(fname, 'r') as y_file: with open(fname, "r") as y_file:
params = yaml.load(y_file, Loader=yaml.FullLoader) params = yaml.load(y_file, Loader=yaml.FullLoader)
logger.info('loaded params...') logger.info("loaded params...")
pp = pprint.PrettyPrinter(indent=4) pp = pprint.PrettyPrinter(indent=4)
pp.pprint(params) pp.pprint(params)
@@ -69,21 +66,24 @@ class Trainer:
def checkpoint(self): def checkpoint(self):
fb_trainer = Trainer(self.fname, True) fb_trainer = Trainer(self.fname, True)
return submitit.helpers.DelayedSubmission(fb_trainer,) return submitit.helpers.DelayedSubmission(
fb_trainer,
)
def launch(): def launch():
executor = submitit.AutoExecutor( executor = submitit.SlurmExecutor(
folder=os.path.join(args.folder, 'job_%j'), folder=os.path.join(args.folder, "job_%j"), max_num_timeout=20
slurm_max_num_timeout=20) )
executor.update_parameters( executor.update_parameters(
slurm_partition=args.partition, partition=args.partition,
slurm_mem_per_gpu='55G', mem_per_gpu="55G",
timeout_min=args.time, time=args.time,
nodes=args.nodes, nodes=args.nodes,
tasks_per_node=args.tasks_per_node, ntasks_per_node=args.tasks_per_node,
cpus_per_task=10, cpus_per_task=10,
gpus_per_node=args.tasks_per_node) gpus_per_node=args.tasks_per_node,
)
config_fnames = [args.fname] config_fnames = [args.fname]
@@ -91,7 +91,9 @@ def launch():
with executor.batch(): with executor.batch():
for cf in config_fnames: for cf in config_fnames:
fb_trainer = Trainer(cf) fb_trainer = Trainer(cf)
job = executor.submit(fb_trainer,) job = executor.submit(
fb_trainer,
)
trainers.append(fb_trainer) trainers.append(fb_trainer)
jobs.append(job) jobs.append(job)
@@ -99,6 +101,6 @@ def launch():
print(job.job_id) print(job.job_id)
if __name__ == '__main__': if __name__ == "__main__":
args = parser.parse_args() args = parser.parse_args()
launch() launch()