From 35634fd2eb3e607daa7e13bc372959cc3dacb73a Mon Sep 17 00:00:00 2001 From: Johannes Fischer Date: Thu, 29 Jul 2021 14:40:19 +0200 Subject: [PATCH] Separate experiment from main.py --- experiments/experiment.py | 141 ++++++++++++++++++++++++++++++++++++++ src/main.py | 140 +------------------------------------ 2 files changed, 142 insertions(+), 139 deletions(-) diff --git a/experiments/experiment.py b/experiments/experiment.py index e69de29..ad79738 100644 --- a/experiments/experiment.py +++ b/experiments/experiment.py @@ -0,0 +1,141 @@ +import json5 +from functools import partial +import os +opj = os.path.join + +from src.main import basestr, main + +def parse_args(): + """ + Parse arguments to main + Returns: + kwargs: dictionary of arguments: + train (bool): whether to run train loop + test (bool): whether to run test loop + method (str): the method to try for imitation + loc (int): the location index of the roundabout + config (str): config path + seed (int): RNG seed + """ + import argparse + parser = argparse.ArgumentParser(description='Save Expert Trajectories') + parser.add_argument('--loc', default=0, type=int, + help='location (default 0)') + parser.add_argument("--train", help="train model", + action="store_true") + parser.add_argument("--ray", help="use ray tune to run multiple experiments", + action="store_true") + parser.add_argument("--test", help="test model", + action="store_true") + parser.add_argument("--method", help="modeling method", + choices=['bc', 'gail', 'advil'], default='bc') + parser.add_argument("--config", help="config file path", + default=None, type=str) + parser.add_argument('--seed', default=0, type=int, + help='seed') + args = parser.parse_args() + kwargs = { + 'train':args.train, + 'test':args.test, + 'method':args.method, + 'loc':args.loc, + 'config_path':args.config, + 'seed':args.seed, + 'ray':args.ray + } + return kwargs + +def get_full_config(ray_config:dict, method:str)->dict: + """ + Get full model configuration from ray config and method string + Args: + ray_config (dict): ray config + method (str): method to get full configuration for + """ + if method == 'bc': + from src.bc import bc_config + config = bc_config(ray_config) + else: + raise NotImplementedError + return config + +def get_ray_config(method:str)->dict: + """ + Get configuration for ray based on method. + Args: + method (str): method to get configuration for + Returns: + ray_config (dict): configuration for ray + """ + if method == 'bc': + ray_config = { + "lr": tune.choice([1e-4, 1e-3, 1e-2, 1e-1]), + "weight_decay": tune.choice([0.001, 0.01, 0.1, 0.5, 0.9]), + "loss": tune.choice(['huber', 'mse']), + "train_batch_size": tune.choice([16,32,64]), + "deepsets_phi_hidden_n": tune.choice([1,2,3]), + "deepsets_phi_hidden_dim": tune.choice([16,32,64]), + "deepsets_latent_dim": tune.choice([16,32,64]), + "deepsets_rho_hidden_n": tune.choice([0,1,2]), + "deepsets_rho_hidden_dim": tune.choice([16,32,64]), + "deepsets_output_dim": tune.choice([8,16,32,64]), + "head_hidden_n": tune.choice([1,2,3]), + "head_hidden_dim": tune.choice([16,32,64]), + "head_final_activation": tune.choice(['sigmoid', None]), + } + else: + raise NotImplementedError + return ray_config + +if __name__ == '__main__': + kwargs = parse_args() + + # make prefix of output files + outdir = opj('output',kwargs['method'],'loc%02i'%(kwargs['loc'])) + + if kwargs['config_path']: + # load config + with open(kwargs['config_path'], 'r') as cfg: + config = json5.load(cfg) + if not os.path.isdir(outdir): + os.makedirs(outdir) + filestr = opj(outdir, basestr(**kwargs)) + 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) + + datadir = os.path.abspath('./expert_data') + ray_config = get_ray_config(kwargs['method']) + custom_scheduler = ASHAScheduler( + metric='cv_loss', + mode="min", + grace_period=25, + ) + analysis = tune.run( + partial(ray_train, datadir=datadir), + config=ray_config, + scheduler=custom_scheduler, + local_dir=outdir, + #resources_per_trial={"cpu": 2}, + time_budget_s=45*60, + num_samples=2, + ) + else: + raise Exception('No valid config found') + + + + + + diff --git a/src/main.py b/src/main.py index aec6622..c51d6df 100644 --- a/src/main.py +++ b/src/main.py @@ -1,12 +1,9 @@ +import os import torch import gym import intersim import numpy as np -import json5 -import os -opj = os.path.join from tqdm import tqdm -from functools import partial from src import InteractionDatasetSingleAgent, metrics from intersim.utils import get_map_path, get_svt @@ -105,138 +102,3 @@ def simulate_policy(policy, loc=0, track=0, filestr='', nframes=float('inf')): pbar.update() env.close(filestr=filestr) - -def parse_args(): - """ - Parse arguments to main - Returns: - kwargs: dictionary of arguments: - train (bool): whether to run train loop - test (bool): whether to run test loop - method (str): the method to try for imitation - loc (int): the location index of the roundabout - config (str): config path - seed (int): RNG seed - """ - import argparse - parser = argparse.ArgumentParser(description='Save Expert Trajectories') - parser.add_argument('--loc', default=0, type=int, - help='location (default 0)') - parser.add_argument("--train", help="train model", - action="store_true") - parser.add_argument("--all-runs", help="use ray tune to run multiple experiments", - action="store_true") - parser.add_argument("--test", help="test model", - action="store_true") - parser.add_argument("--method", help="modeling method", - choices=['bc', 'gail', 'advil'], default='bc') - parser.add_argument("--config", help="config file path", - default=None, type=str) - parser.add_argument('--seed', default=0, type=int, - help='seed') - args = parser.parse_args() - kwargs = { - 'train':args.train, - 'test':args.test, - 'method':args.method, - 'loc':args.loc, - 'config_path':args.config, - 'seed':args.seed, - 'all_runs':args.all_runs - } - return kwargs - -def get_full_config(ray_config:dict, method:str)->dict: - """ - Get full model configuration from ray config and method string - Args: - ray_config (dict): ray config - method (str): method to get full configuration for - """ - if method == 'bc': - from src.bc import bc_config - config = bc_config(ray_config) - else: - raise NotImplementedError - return config - -def get_ray_config(method:str)->dict: - """ - Get configuration for ray based on method. - Args: - method (str): method to get configuration for - Returns: - ray_config (dict): configuration for ray - """ - if method == 'bc': - ray_config = { - "lr": tune.choice([1e-4, 1e-3, 1e-2, 1e-1]), - "weight_decay": tune.choice([0.001, 0.01, 0.1, 0.5, 0.9]), - "loss": tune.choice(['huber', 'mse']), - "train_batch_size": tune.choice([16,32,64]), - "deepsets_phi_hidden_n": tune.choice([1,2,3]), - "deepsets_phi_hidden_dim": tune.choice([16,32,64]), - "deepsets_latent_dim": tune.choice([16,32,64]), - "deepsets_rho_hidden_n": tune.choice([0,1,2]), - "deepsets_rho_hidden_dim": tune.choice([16,32,64]), - "deepsets_output_dim": tune.choice([8,16,32,64]), - "head_hidden_n": tune.choice([1,2,3]), - "head_hidden_dim": tune.choice([16,32,64]), - "head_final_activation": tune.choice(['sigmoid', None]), - } - else: - raise NotImplementedError - return ray_config - -if __name__ == '__main__': - kwargs = parse_args() - - # make prefix of output files - outdir = opj('output',kwargs['method'],'loc%02i'%(kwargs['loc'])) - - if kwargs['config_path']: - # load config - with open(kwargs['config_path'], 'r') as cfg: - config = json5.load(cfg) - if not os.path.isdir(outdir): - os.makedirs(outdir) - filestr = opj(outdir, basestr(**kwargs)) - main(config, filestr=filestr, **kwargs) - - elif kwargs['all_runs'] and kwargs['train']: - - def ray_train(config, datadir=None): - full_config = get_full_config(config, kwargs['method']) - main(full_config, filestr='exp', datadir=datadir, ray=True, **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) - - datadir = os.path.abspath('./expert_data') - ray_config = get_ray_config(kwargs['method']) - custom_scheduler = ASHAScheduler( - metric='cv_loss', - mode="min", - grace_period=25, - ) - analysis = tune.run( - partial(ray_train, datadir=datadir), - config=ray_config, - scheduler=custom_scheduler, - local_dir=outdir, - #resources_per_trial={"cpu": 2}, - time_budget_s=45*60, - num_samples=30, - ) - else: - raise Exception('No valid config found') - - - - - -