adding dtypes and fixing matrix indexing

This commit is contained in:
Arec
2021-07-22 12:06:42 -07:00
parent 1c22bd6111
commit 08eb898812

View File

@@ -14,7 +14,7 @@ class InteractionDatasetMultiAgent(Dataset):
class InteractionDatasetSingleAgent(Dataset): class InteractionDatasetSingleAgent(Dataset):
"""Class to load states and actions for individual agents.""" """Class to load states and actions for individual agents."""
def __init__(self, output_dir='expert_data', loc:int = 0, tracks:list = [0]): def __init__(self, output_dir='expert_data', loc:int = 0, tracks:list = [0], dtype=torch.float32):
""" """
Args: Args:
output_dir (string): Directory with all the images. output_dir (string): Directory with all the images.
@@ -24,6 +24,7 @@ class InteractionDatasetSingleAgent(Dataset):
self.output_dir = output_dir self.output_dir = output_dir
self.loc = loc self.loc = loc
self.tracks = tracks self.tracks = tracks
self.dtype = dtype
self._load_dataset() self._load_dataset()
def _load_dataset(self): def _load_dataset(self):
@@ -43,24 +44,25 @@ class InteractionDatasetSingleAgent(Dataset):
for t in range(T): for t in range(T):
nni = ~torch.isnan(observations[t]['state'][:,0]) nni = ~torch.isnan(observations[t]['state'][:,0])
max_nv = max(max_nv,nni.count_nonzero()) max_nv = max(max_nv,nni.count_nonzero())
self.raw_data['state'].append(observations[t]['state'][nni].float()) self.raw_data['state'].append(observations[t]['state'][nni])
self.raw_data['relative_state'].append(observations[t]['relative_state'][nni.nonzero(),nni.nonzero()].float()) self.raw_data['relative_state'].append(observations[t]['relative_state'].index_select(0,
self.raw_data['action'].append(actions[t][nni].float()) nni.nonzero()[:,0]).index_select(1, nni.nonzero()[:,0]))
self.raw_data['path_x'].append(observations[t]['paths'][0][nni].float()) self.raw_data['action'].append(actions[t][nni])
self.raw_data['path_y'].append(observations[t]['paths'][1][nni].float()) self.raw_data['path_x'].append(observations[t]['paths'][0][nni])
self.raw_data['path_y'].append(observations[t]['paths'][1][nni])
# cat lists # cat lists
self.raw_data['state'] = torch.cat(self.raw_data['state']) self.raw_data['state'] = torch.cat(self.raw_data['state']).type(self.dtype)
self.raw_data['action'] = torch.cat(self.raw_data['action']) self.raw_data['action'] = torch.cat(self.raw_data['action']).type(self.dtype)
self.raw_data['path_x'] = torch.cat(self.raw_data['path_x']) self.raw_data['path_x'] = torch.cat(self.raw_data['path_x']).type(self.dtype)
self.raw_data['path_y'] = torch.cat(self.raw_data['path_y']) self.raw_data['path_y'] = torch.cat(self.raw_data['path_y']).type(self.dtype)
# pad second dimension of relative state # pad second dimension of relative state
for i in range(len(self.raw_data['relative_state'])): for i in range(len(self.raw_data['relative_state'])):
nv1, nv2, d = self.raw_data['relative_state'][i].shape nv1, nv2, d = self.raw_data['relative_state'][i].shape
pad = torch.zeros(nv1, max_nv-nv2, d) * np.nan pad = torch.zeros(nv1, max_nv-nv2, d, dtype=self.dtype) * np.nan
self.raw_data['relative_state'][i] = torch.cat((self.raw_data['relative_state'][i], pad), dim=1) self.raw_data['relative_state'][i] = torch.cat((self.raw_data['relative_state'][i], pad), dim=1)
self.raw_data['relative_state'] = torch.cat(self.raw_data['relative_state']) self.raw_data['relative_state'] = torch.cat(self.raw_data['relative_state']).type(self.dtype)
# mandate equal length # mandate equal length
assert len(self.raw_data['state']) == len(self.raw_data['relative_state']) \ assert len(self.raw_data['state']) == len(self.raw_data['relative_state']) \