From ba79de58b86935b29613d83ee5768c9558654d57 Mon Sep 17 00:00:00 2001 From: Arec Date: Thu, 21 Oct 2021 07:01:12 -0700 Subject: [PATCH] moving feasibility checkers into src.util.collisions, and doing expert processing using the tools in src.data.expert --- .../etienne/intersimple/gail_options_image.py | 110 ++---------------- 1 file changed, 7 insertions(+), 103 deletions(-) diff --git a/scratch/etienne/intersimple/gail_options_image.py b/scratch/etienne/intersimple/gail_options_image.py index 150afe3..71d4741 100644 --- a/scratch/etienne/intersimple/gail_options_image.py +++ b/scratch/etienne/intersimple/gail_options_image.py @@ -1,6 +1,10 @@ # %% +import sys +sys.path.append('../../../') from src.discriminator import CnnDiscriminator, CnnDiscriminatorFlatAction from src.policies import OptionsCnnPolicy +from src.util import feasible +from src.data import load_experts from imitation.algorithms import adversarial from imitation.util import logger @@ -11,7 +15,6 @@ from stable_baselines3.common.env_util import make_vec_env import torch import torch.utils.data -from torch.distributions import Categorical import numpy as np import itertools import gym @@ -21,7 +24,6 @@ import pathlib from tqdm import tqdm from intersim.envs.intersimple import NRasterized, NRasterizedRandomAgent, NRasterizedIncrementingAgent -from intersim.collisions import state_to_polygon ALL_OPTIONS = [(v,t) for v in [0,2,4,6,8] for t in [5, 10, 20]] # option 0 is safe fallback @@ -235,103 +237,6 @@ def generate_plan(env, i): assert len(plan) == t, "incorrect plan length" return plan -def check_future_collisions_circle(env, actions): - """Compute collision information for circular vehicle approximations - - Args: - env (gym.Env): current environment state - actions (list of torch.Tensor): list of B (T, nv, adims) T-length action profiles - Returns: - states (torch.Tensor): tensor of shape (B, T, nv, 5) of future states based on the action profiles - collision_tensor (torch.Tensor): tensor of shape (B, T, nv) of bools indicating which plan collides with which vehicles in which time frame - false: colliding, true: not colliding - """ - B, (T, nv, _) = len(actions), actions[0].shape - - states = torch.stack(env._env.propagate_action_profile(actions), axis=0) - assert states.shape == (B, T, nv, 5) - - distance = ((states[:, :, :, :2] - states[:, :, env._agent:env._agent+1, :2])**2).sum(-1).sqrt() - distance = torch.where(distance.isnan(), np.inf*torch.ones_like(distance), distance) # only collide with spawned agents - distance[:, :, env._agent] = np.inf # cannot collide with itself - assert distance.shape == (B, T, nv) - - radius = (env._env._lengths**2 + env._env._widths**2).sqrt() / 2 - min_distance = radius[env._agent] + radius - min_distance = min_distance.unsqueeze(0).unsqueeze(0) - assert min_distance.shape == (1, 1, nv) - - collision_tensor = distance > min_distance - assert collision_tensor.shape == (B, T, nv) - return states, collision_tensor - -def check_future_collisions_fast(env, actions): - """Checks whether `env._agent` would collide with other agents assuming `actions` as input. - - Vehicles are (over-)approximated by single circles. - - Args: - env (gym.Env): current environment state - actions (list of torch.Tensor): list of B (T, nv, adims) T-length action profiles - Returns: - feasible (torch.Tensor): tensor of shape (B,) indicating whether the respective action profiles are collision-free - """ - _, collision_tensor = check_future_collisions_circle(env, actions) - return collision_tensor.all(-1).all(-1) - -def check_future_collisions_exact(env, actions): - """ - Checks whether `env._agent` would collide with other agents assuming `actions` as input. - - Args: - env (gym.Env): current environment state - actions (list of torch.Tensor): list of B (T, nv, adims) T-length action profiles - Returns: - feasible (torch.Tensor): tensor of shape (B,) indicating whether the respective action profiles are collision-free - """ - # First check with simple circle collision check - states, collision_tensor = check_future_collisions_circle(env, actions) - (B, T, nv, _) = states.shape - # For those that have colliding circles, check exactly - colliding_mask = ~collision_tensor - - ego_states = states[:, :, env._agent:env._agent+1, :].expand(states.shape) - assert ego_states.shape == states.shape - - # get dimensions - lengths = env._env._lengths.expand(states.shape[:3]) - widths = env._env._widths.expand(states.shape[:3]) - ego_lengths = lengths[:, :, env._agent:env._agent+1].expand(lengths.shape) - ego_widths = widths[:, :, env._agent:env._agent+1].expand(widths.shape) - assert lengths.shape == widths.shape == ego_lengths.shape == ego_widths.shape == (B, T, nv) - - # For every collision instance between ego and other vehicle, check whether rectangles intersect - exact_collisions = torch.zeros_like(collision_tensor[colliding_mask]) - for i, (ego_state, ego_length, ego_width, other_state, other_length, other_width) in enumerate(zip( - ego_states[colliding_mask], ego_lengths[colliding_mask], ego_widths[colliding_mask], - states[colliding_mask], lengths[colliding_mask], widths[colliding_mask] - )): - assert ego_state.shape == other_state.shape == (5,) - assert ego_length.shape == ego_width.shape == other_length.shape == other_width.shape == () - p_ego = state_to_polygon(ego_state, ego_length, ego_width) - p_other = state_to_polygon(other_state, other_length, other_width) - exact_collisions[i] = p_ego.intersects(p_other) - - collision_tensor[colliding_mask] = ~exact_collisions - return collision_tensor.all(-1).all(-1) - -def feasible(env, plan, ch): - """Check if input profile is feasible given current `env` state. Action `ch=0` is safe fallback.""" - if ch == 0: - return True - - # zero pad plan - Take (T,) np plan and convert it to (T, nv, 1) torch.Tensor - full_plan = torch.zeros(len(plan), env._env._nv, 1) - full_plan[:, env._agent, 0] = torch.tensor(plan) - # valid = check_future_collisions_fast(env, [full_plan]) # check_future_collisions_fast takes in B-list and outputs (B,) bool tensor - valid = check_future_collisions_exact(env, [full_plan]) # check_future_collisions_fast takes in B-list and outputs (B,) bool tensor - return valid.item() - def flatten_transitions(transitions): return { 'obs': np.stack(list(t['obs'] for t in transitions), axis=0), @@ -426,10 +331,9 @@ if __name__ == '__main__': #env_class = NRasterized #env_settings = {'agent': 51, 'width': 36, 'height': 36, 'm_per_px': 2} + files = ['../../../expert_data/DR_USA_Roundabout_FT0/track%04i/expert.pkl'%(i) for i in range(5)] + transitions=load_experts(files) - with open("data/NormalizedIntersimpleExpertMu.001_NRasterizedAgent51w36h36mppx2.pkl", "rb") as f: - trajectories = pickle.load(f) - transitions = rollout.flatten_trajectories(trajectories) generator = train( transitions, env_class=env_class, @@ -438,7 +342,7 @@ if __name__ == '__main__': discrim_batch_size=32, generator_steps=2048, discount=0.99 - )) + ) generator.save(model_name)