From a15e8c29ff3be17b036a8191e0f61a10ae39dba8 Mon Sep 17 00:00:00 2001 From: Arec Date: Tue, 20 Jul 2021 06:16:54 -0700 Subject: [PATCH] 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=''):