Merge branch 'options-env' into dev

This commit is contained in:
ebuehrle
2021-11-02 09:39:51 +01:00
2 changed files with 153 additions and 23 deletions

View File

@@ -0,0 +1,133 @@
import gym
import torch
from src.util.collisions import feasible
import numpy as np
from collections import deque
import itertools
def imitation_discriminator(discriminator):
return lambda obs, action, next_obs, done: discriminator.discrim_net.predict_reward_train(
state=torch.tensor(obs).unsqueeze(0).to(discriminator.discrim_net.device()),
action=torch.tensor([[action]]).to(discriminator.discrim_net.device()),
next_state=torch.tensor(next_obs).unsqueeze(0).to(discriminator.discrim_net.device()), # unused
done=torch.tensor(done).unsqueeze(0).to(discriminator.discrim_net.device()), # unused
).item()
class OptionsEnv(gym.Wrapper):
def __init__(self, env, options, discriminator, ll_buffer_capacity, *args, **kwargs):
super().__init__(env, *args, **kwargs)
self.options = options
num_hl_options = len(self.options)
self.action_space = gym.spaces.Discrete(num_hl_options)
self.observation_space = gym.spaces.Dict({
'obs': env.observation_space,
'mask': gym.spaces.Box(low=0, high=1, shape=(num_hl_options,)),
})
self.discriminator = discriminator
self.ll_buffer_capacity = ll_buffer_capacity
self.ll_buffer = deque(maxlen=ll_buffer_capacity)
@staticmethod
def _hl_observation(obs, mask):
return {
'obs': obs,
'mask': mask,
}
def reset(self):
self.done = False
self.obs = self.env.reset()
self.m = available_actions(self.env, self.options)
return self._hl_observation(self.obs, self.m)
def _ll_step(self, action):
return self.env.step(action)
def step(self, action):
assert self.m[action]
assert not self.done
plan = list(map(float, generate_plan(self.env, action, self.options)))
reward = 0
steps = 0
while not self.done and plan and \
(feasible(self.env, safety_plan(self.env, plan)) or self.m.sum() == 1):
a, plan = plan[0], plan[1:]
a = self.env._normalize(a)
next_obs, _, self.done, info = self._ll_step(a)
reward += self.discount**steps * self.discriminator(self.obs, a, next_obs, self.done)
self.ll_buffer.append({
'obs': self.obs,
'next_obs': next_obs,
'acts': np.array((a,)),
'dones': np.array(self.done),
})
steps += 1
self.obs = next_obs
self.m = available_actions(self.env, self.options)
return self._hl_observation(self.obs, self.m), reward, self.done, info
def sample_ll(self, n):
assert n <= self.ll_buffer_capacity, f'Sample size of {n} exceeds buffer capacity of {self.ll_buffer_capacity}'
assert n <= len(self.ll_buffer), f'Sample size of {n} exceeds buffer size of {len(self.ll_buffer)}'
return list(itertools.islice(self.ll_buffer, n))
class RenderOptions(OptionsEnv):
def __init__(self, options, *args, **kwargs):
super().__init__(options, discriminator=lambda s, a, n, d: 0, ll_buffer_capacity=0, *args, **kwargs)
def _ll_step(self):
out = super()._ll_step()
self.env.render()
return out
def close(self, *args, **kwargs):
self.env.close(*args, **kwargs)
def safety_plan(env, plan):
return np.concatenate((plan, np.array(5 * [env._env._min_acc])), axis=0)
def available_actions(env, options):
"""Return mask of available actions given current `env` state.
Action 0 is considered safe fallback.
"""
plans = [generate_plan(env, i, options) for i, _ in enumerate(options)]
# is emergency braking still possible?
plans = list(map(lambda p: safety_plan(env, p), plans))
T = max(len(p) for p in plans)
plans = [np.pad(p, ((0, T-len(p)),), constant_values=np.nan) for p in plans]
plans = np.stack(plans, axis=0)
valid = feasible(env, plans)
if not valid.any():
valid[0] = True
return valid
def target_velocity_plan(current_v: float, target_v: float, t: int, dt: float):
"""Smoothly target a velocity in a given number of steps"""
# for now, constant acceleration
a = (target_v - current_v) / (t * dt)
return a*np.ones((t,))
def generate_plan(env, i, options):
"""Generate input profile for high-level action `i`."""
assert i < len(options), "Invalid option index {i}"
target_v, t = options[i]
current_v = env._env.state[env._agent, 1].item() # extract from env
plan = target_velocity_plan(current_v, target_v, t, env._env._dt)
assert len(plan) == t, "incorrect plan length"
return plan

View File

@@ -5,13 +5,7 @@ sys.path.append('../../../')
from src.discriminator import CnnDiscriminatorFlatAction from src.discriminator import CnnDiscriminatorFlatAction
from imitation.algorithms import adversarial from imitation.algorithms import adversarial
import stable_baselines3 import stable_baselines3
import torch.utils.data
import numpy as np
from intersim.envs import NRasterizedRouteSpeedRandomAgentLocation from intersim.envs import NRasterizedRouteSpeedRandomAgentLocation
import itertools
from torch.distributions import Categorical
import gym
import torch
import pickle import pickle
import imitation.data.rollout as rollout import imitation.data.rollout as rollout
import tempfile import tempfile
@@ -20,8 +14,9 @@ from imitation.util import logger
from stable_baselines3.common.env_util import make_vec_env from stable_baselines3.common.env_util import make_vec_env
from tqdm import tqdm from tqdm import tqdm
from src.policies.options import OptionsCnnPolicy from src.policies.options import OptionsCnnPolicy
from src.gail.options import OptionsEnv, LLOptions, HLOptions, RenderOptions from src.gail.train import flatten_transitions
from src.gail.train import train_discriminator, train_generator
from gail.options2 import OptionsEnv, RenderOptions, imitation_discriminator
model_name = 'gail_options_image_random_location' model_name = 'gail_options_image_random_location'
env_settings = {'width': 70, 'height': 70, 'm_per_px': 1, 'map_color': 128, 'mu': 0.001} env_settings = {'width': 70, 'height': 70, 'm_per_px': 1, 'map_color': 128, 'mu': 0.001}
@@ -30,12 +25,13 @@ ALL_OPTIONS = [(v,t) for v in [0,2,4,6,8] for t in [5, 10, 20]] # option 0 is sa
def train( def train(
expert_data, expert_data,
epochs=200,
expert_batch_size=1024, expert_batch_size=1024,
generator_steps=1024, discriminator_updates_per_round=10,
generator_steps=256,
generator_total_steps=1024,
generator_updates_per_round=10,
discount=0.99, discount=0.99,
n_disc_updates_per_round=10, epochs=200,
n_gen_updates_per_round=10,
): ):
env = NRasterizedRouteSpeedRandomAgentLocation(**env_settings) env = NRasterizedRouteSpeedRandomAgentLocation(**env_settings)
env.discount = discount env.discount = discount
@@ -55,24 +51,25 @@ def train(
gen_algo=stable_baselines3.PPO("CnnPolicy", venv), # unused gen_algo=stable_baselines3.PPO("CnnPolicy", venv), # unused
) )
options_env = OptionsEnv(env, discriminator=imitation_discriminator(discriminator), options=ALL_OPTIONS, ll_buffer_capacity=expert_batch_size)
generator = stable_baselines3.PPO( generator = stable_baselines3.PPO(
OptionsCnnPolicy, OptionsCnnPolicy,
OptionsEnv(env, options=ALL_OPTIONS), options_env,
verbose=1, verbose=1,
n_steps=generator_steps, n_steps=generator_steps,
n_epochs=n_gen_updates_per_round, n_epochs=generator_updates_per_round,
)
# PPO.train requires logger as set up in
# PPO._setup_learn (called by PPO.learn)
generator._logger = stable_baselines3.common.utils.configure_logger(
generator.verbose,
generator.tensorboard_log,
) )
for _ in tqdm(range(epochs)): for _ in tqdm(range(epochs)):
train_discriminator(LLOptions(env, options=ALL_OPTIONS), generator, discriminator, num_samples=expert_batch_size, n_updates=n_disc_updates_per_round) # train generator
train_generator(HLOptions(env, options=ALL_OPTIONS), generator, discriminator, num_samples=generator_steps) generator.learn(total_timesteps=generator_total_steps)
# train discriminator
generator_samples = options_env.sample_ll(expert_batch_size)
generator_samples = flatten_transitions(generator_samples)
for _ in range(discriminator_updates_per_round):
discriminator.train_disc(gen_samples=generator_samples)
generator.save(model_name) generator.save(model_name)
return generator return generator