making test case for typing bug and fixing some small typing errors in bc

This commit is contained in:
Arec
2021-07-22 08:23:46 -07:00
parent 827a8e7172
commit 91d052445e
6 changed files with 42 additions and 17 deletions

View File

@@ -43,11 +43,11 @@ class InteractionDatasetSingleAgent(Dataset):
for t in range(T):
nni = ~torch.isnan(observations[t]['state'][:,0])
max_nv = max(max_nv,nni.count_nonzero())
self.raw_data['state'].append(observations[t]['state'][nni])
self.raw_data['relative_state'].append(observations[t]['relative_state'][nni.nonzero(),nni.nonzero()])
self.raw_data['action'].append(actions[t][nni])
self.raw_data['path_x'].append(observations[t]['paths'][0][nni])
self.raw_data['path_y'].append(observations[t]['paths'][1][nni])
self.raw_data['state'].append(observations[t]['state'][nni].float())
self.raw_data['relative_state'].append(observations[t]['relative_state'][nni.nonzero(),nni.nonzero()].float())
self.raw_data['action'].append(actions[t][nni].float())
self.raw_data['path_x'].append(observations[t]['paths'][0][nni].float())
self.raw_data['path_y'].append(observations[t]['paths'][1][nni].float())
# cat lists
self.raw_data['state'] = torch.cat(self.raw_data['state'])