Merge branch 'main' of github.com:sisl/InteractionImitation

This commit is contained in:
Johannes Fischer
2021-07-19 17:28:49 +02:00

View File

@@ -6,7 +6,7 @@ import numpy as np
import intersim import intersim
from intersim.utils import get_map_path, get_svt, SVT_to_sim_stateactions from intersim.utils import get_map_path, get_svt, SVT_to_sim_stateactions
from intersim import collisions
import os import os
opj = os.path.join opj = os.path.join
@@ -36,15 +36,29 @@ def generate_expert_data(path: str='expert_data', loc: int = 0, track:int = 0, *
env.reset() env.reset()
done = False done = False
obs, actions_taken = [], [] obs, actions_taken, max_devs = [], [], []
i = 0 i = 0
print(actions.shape) while not done and i < len(actions):
while not done: # check state deviation
env_state = env.projected_state
nni = ~torch.isnan(env_state[:,0])
norms = torch.norm(env_state[nni,:2]-states[i,nni,:2], dim=1)
max_devs.append(norms.max())
# print("Step: %04i, Maximum Deviation: %f m" %(i, max_devs[-1]))
# propagate environment
ob, r, done, info = env.step(actions[i]) ob, r, done, info = env.step(actions[i])
obs.append(ob) obs.append(ob)
actions_taken.append(info['action_taken']) actions_taken.append(info['action_taken'])
i += 1 i += 1
print("Maximum environment deviation from track: %f m" %(max(max_devs)))
# check for collisions
x = torch.stack([ob['state'] for ob in obs])
cols = collisions.check_collisions_trajectory(x, svt.lengths, svt.widths)
assert ~torch.any(cols), 'Error: Collisions found at indices {}'.format(cols.nonzero(as_tuple=True))
# shift actions # shift actions
actions_taken.pop(0) actions_taken.pop(0)
obs.pop(-1) obs.pop(-1)