removing outliers from expert tracks
This commit is contained in:
@@ -69,15 +69,18 @@ def generate_expert_data(path: str='expert_data', loc: int = 0, track:int = 0, *
|
|||||||
torch.save(actions, filestr+'_raw_actions.pt')
|
torch.save(actions, filestr+'_raw_actions.pt')
|
||||||
process_expert_observations(obs, actions, filestr)
|
process_expert_observations(obs, actions, filestr)
|
||||||
|
|
||||||
def process_expert_observations(obs, actions, filestr, dtype=torch.float32):
|
def process_expert_observations(obs, actions, filestr, remove_outliers=True, dtype=torch.float32):
|
||||||
"""
|
"""
|
||||||
Process the expert observations and save them as torch tensors
|
Process the expert observations and save them as torch tensors
|
||||||
Args:
|
Args:
|
||||||
obs (list[dict]): lost of observations
|
obs (list[dict]): lost of observations
|
||||||
actions (torch.Tensor): (T, nv, a) tensor of actions
|
actions (torch.Tensor): (T, nv, a) tensor of actions
|
||||||
filestr (str): base filename with which to save out observation tensors
|
filestr (str): base filename with which to save out observation tensors
|
||||||
|
remove_outliers (bool): whether to remove datapoints with acceleration above or below 5 m/s/s
|
||||||
|
dtype (torch.Type): type to convert data to
|
||||||
"""
|
"""
|
||||||
data = {'state':[], 'action':[], 'relative_state':[], 'path_x':[], 'path_y':[]}
|
keys = ['state', 'action', 'relative_state', 'path_x', 'path_y']
|
||||||
|
data = {key:[] for key in keys}
|
||||||
assert len(obs) == len(actions), 'non-matching action and observation lengths'
|
assert len(obs) == len(actions), 'non-matching action and observation lengths'
|
||||||
T = len(obs)
|
T = len(obs)
|
||||||
max_nv = 0
|
max_nv = 0
|
||||||
@@ -104,6 +107,11 @@ def process_expert_observations(obs, actions, filestr, dtype=torch.float32):
|
|||||||
data['relative_state'][i] = torch.cat((data['relative_state'][i], pad), dim=1)
|
data['relative_state'][i] = torch.cat((data['relative_state'][i], pad), dim=1)
|
||||||
data['relative_state'] = torch.cat(data['relative_state']).type(dtype)
|
data['relative_state'] = torch.cat(data['relative_state']).type(dtype)
|
||||||
|
|
||||||
|
if remove_outliers:
|
||||||
|
non_outlier_indices = torch.nonzero(torch.abs(data['action'][:,0]) < 5)
|
||||||
|
for key in keys:
|
||||||
|
data[key] = data[key][non_outlier_indices[:,0]]
|
||||||
|
|
||||||
# mandate equal length
|
# mandate equal length
|
||||||
assert len(data['state']) == len(data['relative_state']) \
|
assert len(data['state']) == len(data['relative_state']) \
|
||||||
== len(data['action']) == len(data['path_x']) \
|
== len(data['action']) == len(data['path_x']) \
|
||||||
|
|||||||
@@ -51,8 +51,6 @@ def visualize_distribution(true, pred, filestr):
|
|||||||
"""
|
"""
|
||||||
nni1 = ~torch.isnan(true)
|
nni1 = ~torch.isnan(true)
|
||||||
nni2 = ~torch.isnan(pred)
|
nni2 = ~torch.isnan(pred)
|
||||||
import pdb
|
|
||||||
pdb.set_trace()
|
|
||||||
plt.figure()
|
plt.figure()
|
||||||
plt.hist(true[nni1].numpy(), density=True, bins=20)
|
plt.hist(true[nni1].numpy(), density=True, bins=20)
|
||||||
plt.hist(pred[nni2].numpy(), density=True, bins=20)
|
plt.hist(pred[nni2].numpy(), density=True, bins=20)
|
||||||
|
|||||||
Reference in New Issue
Block a user