From 2e29f42d7579d7bd8b28af63e757c5e2635f5483 Mon Sep 17 00:00:00 2001 From: Arec Date: Tue, 20 Jul 2021 05:46:43 -0700 Subject: [PATCH 1/3] add flag to process all tracks of a particular location --- src/expert_data.py | 8 +++++++- 1 file changed, 7 insertions(+), 1 deletion(-) diff --git a/src/expert_data.py b/src/expert_data.py index 70921ce..0b47c35 100644 --- a/src/expert_data.py +++ b/src/expert_data.py @@ -91,5 +91,11 @@ if __name__ == '__main__': help='location (default 0)') parser.add_argument('--track', default=0, type=int, help='track number (default 0)') + parser.add_argument('--all-tracks', action='store_true', + help='whether to process all tracks at location') args = parser.parse_args() - generate_expert_data(loc=args.loc,track=args.track) \ No newline at end of file + if args.all_tracks: + for i in range(intersim.MAX_TRACKS): + generate_expert_data(loc=args.loc, track=i) + else: + generate_expert_data(loc=args.loc,track=args.track) \ No newline at end of file From 2fb5d5e5b15c56dd66623308641ee6d5e0a8acab Mon Sep 17 00:00:00 2001 From: Arec Date: Tue, 20 Jul 2021 05:58:24 -0700 Subject: [PATCH 2/3] developing main experiment loop, functions required to implement in bc and other imitation methods --- src/__init__.py | 3 +- src/bc/__init__.py | 1 + src/bc/bc.py | 13 +++++ src/main.py | 127 +++++++++++++++++++++++++++++++++++++++++++++ 4 files changed, 143 insertions(+), 1 deletion(-) create mode 100644 src/bc/bc.py create mode 100644 src/main.py diff --git a/src/__init__.py b/src/__init__.py index 739d996..db75ef3 100644 --- a/src/__init__.py +++ b/src/__init__.py @@ -1 +1,2 @@ -from src.expert_data import generate_expert_data, load_expert_data \ No newline at end of file +from src.expert_data import generate_expert_data, load_expert_data +from src.data_utils import InteractionDatasetSingleAgent \ No newline at end of file diff --git a/src/bc/__init__.py b/src/bc/__init__.py index e69de29..2b0a673 100644 --- a/src/bc/__init__.py +++ b/src/bc/__init__.py @@ -0,0 +1 @@ +from src.bc.bc import BehaviorCloningPolicy, train, load_policy, metrics diff --git a/src/bc/bc.py b/src/bc/bc.py new file mode 100644 index 0000000..5b9a368 --- /dev/null +++ b/src/bc/bc.py @@ -0,0 +1,13 @@ +class BehaviorCloningPolicy(): + pass + +def load_policy(): + pass + +def metrics(): + pass + +def train(): + pass + + diff --git a/src/main.py b/src/main.py new file mode 100644 index 0000000..5dc3f64 --- /dev/null +++ b/src/main.py @@ -0,0 +1,127 @@ +import torch +import gym +import intersim +import numpy as np +import os +opj = os.path.join +from src import InteractionDatasetSingleAgent + +def basestr(**kwargs): + """ + Return base prefix for all files relating to a certain experiment + Args: + kwargs (dict): keyword arguments sent to main training loop + Returns: + basestr (str): prefix + """ + return 'base_' + +def main(method='bc', train=False, test=False, loc=0, **kwargs): + """ + Main loop for training and testing different imitation models + Args: + 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 + kwargs (dict): remaining kwargs for policy and training loop + """ + outdir = opj('output',method,'loc%02i'%(loc)) + if not os.path.isdir(outdir): + os.mkdir(outdir) + filestr = opj(outdir, basestr(**kwargs)) + + # define transforms + transforms={} + + # method-based training + if method=='bc': + from src import bc + policy = bc.BehaviorCloningPolicy(transforms=transforms, **kwargs) + load_policy = bc.load_policy + metrics = bc.metrics + train = bc.train + else: + raise NotImplementedError + + # default train / cv / test split datasets + if train: + train_dataset = InteractionDatasetSingleAgent(loc=loc, tracks=[0,1,2], transforms=transforms, train=True) + cv_dataset = InteractionDatasetSingleAgent(loc=loc, tracks=[3], transforms=transforms) + train(train_dataset, cv_dataset, policy, filestr=filestr, **kwargs) + + if test: + + # load policy + policy = load_policy(filestr=filestr) + + # simulate policy + test_track = 4 + simulate_policy(policy, loc=loc, track=track, filestr=filestr) + + # run test metrics + # test_dataset = InteractionDatasetSingleAgent(loc=loc, tracks=[4], transforms=transforms, train=False) + # metrics(test_datset, policy) + + +def simulate_policy(policy, loc=0, track=0, filestr=''): + """ + Simulate a trained policy + Args: + policy: the policy to simulate, which should return action directly + loc (int): location index to test policy + track (int): track to test policy + filestr (str): path prefix to save simulation to + """ + # animate from environment + env = gym.make('intersim:intersim-v0', loc=loc, track=track, + min_acc=-np.inf, max_acc=np.inf) + + ob, _ = env.reset() + env.render() + done = False + while not done: + # get action + action = policy(ob) + + # propagate environment + ob, r, done, info = env.step(action) + env.render() + + 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 + """ + 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("--test", help="test model", + action="store_true") + parser.add_argument("--method", help="modeling method", + choices=['bc', 'gail', 'advil'], default='bc') + parser.add_argument() + parser.add_argument() + args = parser.parse_args() + kwargs = { + 'train'=args.train, + 'test'=args.test, + 'method'=args.method, + 'loc'=args.loc + } + return kwargs + + +if __name__ == '__main__': + kwargs = parse_args() + main(**kwargs) \ No newline at end of file From a15e8c29ff3be17b036a8191e0f61a10ae39dba8 Mon Sep 17 00:00:00 2001 From: Arec Date: Tue, 20 Jul 2021 06:16:54 -0700 Subject: [PATCH 3/3] moving transforms out of dataset class, will be exclusively in policy classes --- src/data_utils.py | 18 ++---------------- src/main.py | 20 ++++++++++---------- 2 files changed, 12 insertions(+), 26 deletions(-) diff --git a/src/data_utils.py b/src/data_utils.py index aba6c13..dbf41f8 100644 --- a/src/data_utils.py +++ b/src/data_utils.py @@ -1,7 +1,6 @@ import torch -from torch.utils.data import Dataset, DataLoader +from torch.utils.data import Dataset import numpy as np -#from torchvision import transforms, utils from src.expert_data import load_expert_data import os opj = os.path.join @@ -15,24 +14,16 @@ class InteractionDatasetMultiAgent(Dataset): class InteractionDatasetSingleAgent(Dataset): """Class to load states and actions for individual agents.""" - def __init__(self, output_dir='expert_data', loc:int = 0, tracks:list = [0], transforms={}): + def __init__(self, output_dir='expert_data', loc:int = 0, tracks:list = [0]): """ Args: output_dir (string): Directory with all the images. loc (int): location index tracks (list[int]): track indices - transforms (dict): dictionary of transforms to apply to different variables """ self.output_dir = output_dir self.loc = loc self.tracks = tracks - self.transforms = transforms - #self.action_transform = transforms.get('action', None) - #self.state_transform = transforms.get('state', None) - #self.relative_state_transform = transforms.get('relative_state', None) - #self.paths_x_transform = transforms.get('paths_x', None) - #self.paths_y_transform = transform.get('paths_y',None) - self._load_dataset() def _load_dataset(self): @@ -95,9 +86,4 @@ class InteractionDatasetSingleAgent(Dataset): """ keys = ['state', 'relative_state', 'path_x', 'path_y', 'action'] sample = {key:self.raw_data[key][idx] for key in keys} - - for key in keys: - if key in self.transforms.keys(): - sample[key] = self.transforms[key](sample[key]) - return sample \ No newline at end of file diff --git a/src/main.py b/src/main.py index 5dc3f64..c071dc4 100644 --- a/src/main.py +++ b/src/main.py @@ -37,31 +37,31 @@ def main(method='bc', train=False, test=False, loc=0, **kwargs): # method-based training if method=='bc': from src import bc - policy = bc.BehaviorCloningPolicy(transforms=transforms, **kwargs) - load_policy = bc.load_policy - metrics = bc.metrics - train = bc.train + policy_class = bc.BehaviorCloningPolicy + load_policy_fn = bc.load_policy + metrics_fn = bc.metrics + train_fn = bc.train else: raise NotImplementedError # default train / cv / test split datasets if train: - train_dataset = InteractionDatasetSingleAgent(loc=loc, tracks=[0,1,2], transforms=transforms, train=True) - cv_dataset = InteractionDatasetSingleAgent(loc=loc, tracks=[3], transforms=transforms) - train(train_dataset, cv_dataset, policy, filestr=filestr, **kwargs) + train_dataset = InteractionDatasetSingleAgent(loc=loc, tracks=[0,1,2]) + cv_dataset = InteractionDatasetSingleAgent(loc=loc, tracks=[3]) + train_fn(train_dataset, cv_dataset, policy_class, filestr=filestr, **kwargs) if test: # load policy - policy = load_policy(filestr=filestr) + policy = load_policy_fn(filestr=filestr) # simulate policy test_track = 4 simulate_policy(policy, loc=loc, track=track, filestr=filestr) # run test metrics - # test_dataset = InteractionDatasetSingleAgent(loc=loc, tracks=[4], transforms=transforms, train=False) - # metrics(test_datset, policy) + # test_dataset = InteractionDatasetSingleAgent(loc=loc, tracks=[4]) + # metrics_fn(test_dataset, policy) def simulate_policy(policy, loc=0, track=0, filestr=''):