From 3b623a14675fd72e7a0aff24996ab38245569bd0 Mon Sep 17 00:00:00 2001 From: Johannes Fischer Date: Thu, 29 Jul 2021 15:30:41 +0200 Subject: [PATCH] Add testing script for raytune experiments --- experiments/experiment.py | 18 +++++++++++++----- src/main.py | 2 +- 2 files changed, 14 insertions(+), 6 deletions(-) diff --git a/experiments/experiment.py b/experiments/experiment.py index ad79738..593f6b0 100644 --- a/experiments/experiment.py +++ b/experiments/experiment.py @@ -103,17 +103,17 @@ if __name__ == '__main__': main(config, filestr=filestr, **kwargs) elif kwargs['ray'] and kwargs['train']: - - def ray_train(config, datadir=None): - full_config = get_full_config(config, kwargs['method']) - main(full_config, filestr='exp', datadir=datadir, **kwargs) - # set up ray tune import ray from ray import tune from ray.tune.schedulers import ASHAScheduler + ray.shutdown() ray.init(log_to_driver=False) + + def ray_train(config, datadir=None): + full_config = get_full_config(config, kwargs['method']) + main(full_config, filestr='exp', datadir=datadir, **kwargs) datadir = os.path.abspath('./expert_data') ray_config = get_ray_config(kwargs['method']) @@ -131,6 +131,14 @@ if __name__ == '__main__': time_budget_s=45*60, num_samples=2, ) + elif kwargs['ray'] and kwargs['test']: + import ray + from ray.tune import Analysis, ExperimentAnalysis + analysis = Analysis(outdir, default_metric="cv_loss", default_mode="min") + config = analysis.get_best_config() + filepath = analysis.get_best_logdir() + print(filepath) + main(None, filestr=opj(filepath, 'exp'), **kwargs) else: raise Exception('No valid config found') diff --git a/src/main.py b/src/main.py index c51d6df..76e83c7 100644 --- a/src/main.py +++ b/src/main.py @@ -27,7 +27,7 @@ def main(config, method='bc', train=False, test=False, loc=0, datadir='./expert_ test (bool): whether to run test loop method (str): the method to try for imitation loc (int): the location index of the roundabout - config_file (str): path to config file + datadir (str): path to expert data kwargs (dict): remaining kwargs for training loop """ # get/set seed