Add SHAIL
This commit is contained in:
@@ -1,10 +1,10 @@
|
||||
import torch
|
||||
import torch.nn.functional as F
|
||||
from dataclasses import dataclass
|
||||
from core.reparam_module import ReparamPolicy
|
||||
from core.sampling import rollout
|
||||
from core.trpo import trpo_step
|
||||
from core.ppo import ppo_step
|
||||
from src.core.reparam_module import ReparamPolicy
|
||||
from src.core.sampling import rollout
|
||||
from src.core.trpo import trpo_step
|
||||
from src.core.ppo import ppo_step
|
||||
from tqdm import tqdm
|
||||
|
||||
class TerminalLogger:
|
||||
|
||||
@@ -1,6 +1,6 @@
|
||||
import torch
|
||||
from core.sampling import rollout
|
||||
from core.value_estimation import gae
|
||||
from src.core.sampling import rollout
|
||||
from src.core.value_estimation import gae
|
||||
|
||||
def ppo(env_fn, value, policy, epochs, rollout_episodes, rollout_steps, gamma, gae_lambda, clip_ratio, pi_opt, pi_iters, v_opt, v_iters, target_kl=None, max_grad_norm=None):
|
||||
|
||||
|
||||
@@ -1,8 +1,8 @@
|
||||
import torch
|
||||
from core.reparam_module import ReparamPolicy
|
||||
from core.sampling import rollout
|
||||
from core.value_estimation import gae
|
||||
from core.optimization import conjugate_gradient, line_search
|
||||
from src.core.reparam_module import ReparamPolicy
|
||||
from src.core.sampling import rollout
|
||||
from src.core.value_estimation import gae
|
||||
from src.core.optimization import conjugate_gradient, line_search
|
||||
|
||||
def trpo(env_fn, value, policy, epochs, rollout_episodes, rollout_steps, gamma, gae_lambda, delta, backtrack_coeff, backtrack_iters, v_opt, v_iters, cg_iters=10, cg_damping=0.1):
|
||||
|
||||
|
||||
@@ -11,6 +11,8 @@ from src.evaluation.metrics import divergence, visualize_distribution
|
||||
from src.core.policy import SetPolicy, SetDiscretePolicy
|
||||
from src.core.reparam_module import ReparamPolicy
|
||||
from src.options import envs as options_envs2
|
||||
from src.safe_options.policy import SetMaskedDiscretePolicy
|
||||
from src.safe_options import options as options_envs3
|
||||
|
||||
from typing import Optional, List, Dict, Tuple
|
||||
import torch
|
||||
@@ -62,8 +64,14 @@ def load_policy(method:str,
|
||||
policy.load_state_dict(torch.load(policy_file))
|
||||
policy.eval()
|
||||
elif method == 'sgail':
|
||||
policy = sb3.PPO.load(policy_file)
|
||||
raise NotImplementedError
|
||||
policy = SetMaskedDiscretePolicy(env.action_space.n)
|
||||
policy(
|
||||
torch.zeros(env.observation_space['observation'].shape),
|
||||
torch.zeros(env.observation_space['safe_actions'].shape)
|
||||
)
|
||||
policy = ReparamPolicy(policy)
|
||||
policy.load_state_dict(torch.load(policy_file))
|
||||
policy.eval()
|
||||
else:
|
||||
raise NotImplementedError
|
||||
return policy
|
||||
@@ -181,6 +189,7 @@ def evaluate_policy(locations:List[Tuple[int,int]],
|
||||
envs_dict = dict(intersim.envs.intersimple.__dict__)
|
||||
envs_dict.update(dict(options_envs.__dict__))
|
||||
envs_dict.update(dict(options_envs2.__dict__))
|
||||
envs_dict.update(dict(options_envs3.__dict__))
|
||||
policy_metrics = [None]* len(locations)
|
||||
|
||||
# iterate through vehicles
|
||||
|
||||
@@ -7,6 +7,7 @@ import os
|
||||
import pickle
|
||||
from tqdm import tqdm
|
||||
from src.options.envs import OptionsEnv
|
||||
from src.util.wrappers import OptionsTimeLimit
|
||||
|
||||
class IntersimpleEvaluation:
|
||||
"""
|
||||
@@ -35,7 +36,7 @@ class IntersimpleEvaluation:
|
||||
self.env = eval_env
|
||||
self.n_episodes = eval_env.nv
|
||||
self.use_pbar = use_pbar
|
||||
self.is_options_env = isinstance(self.env, OptionsEnv)
|
||||
self.is_options_env = isinstance(self.env, (OptionsEnv, OptionsTimeLimit))
|
||||
|
||||
# metrics present on every step of every episode
|
||||
self.metric_keys_all = ['v_all', 'a_all', 'col_all']
|
||||
|
||||
@@ -3,14 +3,18 @@ import numpy as np
|
||||
import torch
|
||||
from stable_baselines3.common.vec_env import DummyVecEnv as VecEnv
|
||||
|
||||
from core.reparam_module import ReparamPolicy
|
||||
from src.core.reparam_module import ReparamPolicy
|
||||
from tqdm import tqdm
|
||||
from core.gail import train_discriminator, roll_buffer, TerminalLogger
|
||||
from src.core.gail import train_discriminator, roll_buffer, TerminalLogger
|
||||
from dataclasses import dataclass
|
||||
from safe_options.policy_gradient import trpo_step, ppo_step
|
||||
from src.safe_options.policy_gradient import trpo_step, ppo_step
|
||||
import torch.nn.functional as F
|
||||
|
||||
from safe_options.collisions import feasible
|
||||
from src.options.envs import OptionsEnv
|
||||
from src.safe_options.collisions import feasible
|
||||
|
||||
from intersim.envs import IntersimpleLidarFlatIncrementingAgent
|
||||
from src.util.wrappers import OptionsTimeLimit, Setobs, TransformObservation
|
||||
|
||||
@dataclass
|
||||
class Buffer:
|
||||
@@ -156,81 +160,6 @@ def rollout(env_fn, policy, n_episodes, max_steps_per_episode):
|
||||
|
||||
return (states, safe_actions, actions, rewards, dones), (ll_states, ll_actions, ll_rewards, ll_dones)
|
||||
|
||||
class OptionsEnv(gym.Wrapper):
|
||||
|
||||
def __init__(self, env, options):
|
||||
super().__init__(env)
|
||||
self.ll_action_space = env.action_space
|
||||
self.options = options
|
||||
self.action_space = gym.spaces.Discrete(len(options))
|
||||
self.max_plan_length = max(t for _, t in options)
|
||||
|
||||
def plan(self, option):
|
||||
target_v, t = option
|
||||
current_v = self.env._env.state[self.env._agent, 1].item()
|
||||
dt = self.env._env._dt
|
||||
a = (target_v - current_v) / (t * dt)
|
||||
a = self.env._normalize(a)
|
||||
a = a * np.ones((t,))
|
||||
a += 0.01 * np.random.randn(*a.shape)
|
||||
a = np.clip(a, self.ll_action_space.low, self.ll_action_space.high)
|
||||
return a
|
||||
|
||||
def execute_plan(self, obs, option, render_mode=None):
|
||||
observations = np.zeros((self.max_plan_length + 1, *self.env.observation_space.shape))
|
||||
actions = np.zeros((self.max_plan_length + 1, *self.ll_action_space.shape))
|
||||
rewards = np.zeros((self.max_plan_length + 1,))
|
||||
env_done = np.ones((self.max_plan_length + 1,), dtype=bool)
|
||||
plan_done = np.ones((self.max_plan_length + 1,), dtype=bool)
|
||||
infos = []
|
||||
|
||||
plan = self.plan(option)
|
||||
observations[0] = obs
|
||||
env_done[0] = False
|
||||
for k, u in enumerate(plan):
|
||||
plan_done[k] = False
|
||||
o, r, d, i = self.env.step(u)
|
||||
actions[k] = u
|
||||
rewards[k] = r
|
||||
env_done[k+1] = d
|
||||
infos.append(i)
|
||||
observations[k+1] = o
|
||||
|
||||
if render_mode is not None:
|
||||
self.env.render(render_mode)
|
||||
|
||||
if d:
|
||||
break
|
||||
|
||||
n_steps = k + 1
|
||||
return observations, actions, rewards, env_done, plan_done, infos, n_steps
|
||||
|
||||
def step(self, action, render_mode=None):
|
||||
a = int(action)
|
||||
assert a == action
|
||||
ll_obs, ll_actions, ll_rewards, ll_env_done, ll_plan_done, ll_infos, ll_steps = self.execute_plan(self.last_obs, self.options[a], render_mode)
|
||||
hl_obs = ll_obs[ll_steps]
|
||||
hl_reward = (ll_rewards * ~ll_plan_done).sum().item()
|
||||
hl_done = ll_env_done[ll_steps].item()
|
||||
hl_infos = {
|
||||
'll': {
|
||||
'observations': ll_obs,
|
||||
'actions': ll_actions,
|
||||
'rewards': ll_rewards,
|
||||
'env_done': ll_env_done,
|
||||
'plan_done': ll_plan_done,
|
||||
'infos': ll_infos,
|
||||
'steps': ll_steps,
|
||||
}
|
||||
}
|
||||
self.last_obs = hl_obs
|
||||
return hl_obs, hl_reward, hl_done, hl_infos
|
||||
|
||||
def reset(self, *args, **kwargs):
|
||||
self.last_obs = super().reset(*args, **kwargs)
|
||||
return self.last_obs
|
||||
|
||||
|
||||
class SafeOptionsEnv(OptionsEnv):
|
||||
|
||||
def __init__(self, env, options, safe_actions_collision_method=None, abort_unsafe_collision_method=None):
|
||||
@@ -287,7 +216,7 @@ class SafeOptionsEnv(OptionsEnv):
|
||||
o, r, d, i = self.env.step(u)
|
||||
actions[k] = u
|
||||
rewards[k] = r
|
||||
env_done[k+1] = d
|
||||
env_done[k] = d
|
||||
infos.append(i)
|
||||
observations[k+1] = o
|
||||
|
||||
@@ -303,3 +232,29 @@ class SafeOptionsEnv(OptionsEnv):
|
||||
|
||||
n_steps = k + 1
|
||||
return observations, actions, rewards, env_done, plan_done, infos, n_steps
|
||||
|
||||
obs_min = np.array([
|
||||
[-1000, -1000, 0, -np.pi, -1e-1, 0.],
|
||||
[0, -np.pi, -20, -20, -np.pi, -1e-1],
|
||||
[0, -np.pi, -20, -20, -np.pi, -1e-1],
|
||||
[0, -np.pi, -20, -20, -np.pi, -1e-1],
|
||||
[0, -np.pi, -20, -20, -np.pi, -1e-1],
|
||||
[0, -np.pi, -20, -20, -np.pi, -1e-1],
|
||||
]).reshape(-1)
|
||||
|
||||
obs_max = np.array([
|
||||
[1000, 1000, 20, np.pi, 1e-1, 0.],
|
||||
[50, np.pi, 20, 20, np.pi, 1e-1],
|
||||
[50, np.pi, 20, 20, np.pi, 1e-1],
|
||||
[50, np.pi, 20, 20, np.pi, 1e-1],
|
||||
[50, np.pi, 20, 20, np.pi, 1e-1],
|
||||
[50, np.pi, 20, 20, np.pi, 1e-1],
|
||||
]).reshape(-1)
|
||||
|
||||
def NormalizedSafeOptionsEvalEnv(max_episode_steps=float('inf'), safe_actions_collision_method=None, abort_unsafe_collision_method=None, **kwargs):
|
||||
return OptionsTimeLimit(SafeOptionsEnv(Setobs(
|
||||
TransformObservation(IntersimpleLidarFlatIncrementingAgent(
|
||||
n_rays=5,
|
||||
**kwargs,
|
||||
), lambda obs: (obs - obs_min) / (obs_max - obs_min + 1e-10))
|
||||
), options=[(0, 5), (1, 5), (2, 5), (4, 5), (6, 5), (8, 5)], safe_actions_collision_method=safe_actions_collision_method, abort_unsafe_collision_method=abort_unsafe_collision_method), max_episode_steps=max_episode_steps)
|
||||
|
||||
@@ -2,7 +2,7 @@ import torch
|
||||
import torch.nn as nn
|
||||
from torch.distributions import Categorical
|
||||
from torch.distributions.kl import kl_divergence
|
||||
from core.policy import SetDiscretePolicy
|
||||
from src.core.policy import SetDiscretePolicy
|
||||
|
||||
class SetMaskedDiscretePolicy(SetDiscretePolicy):
|
||||
|
||||
@@ -21,6 +21,15 @@ class SetMaskedDiscretePolicy(SetDiscretePolicy):
|
||||
a = super().torch_dist(logits).probs
|
||||
return (a * (1 - z)).sum(-1)
|
||||
|
||||
def predict(self, observations, state=None, episode_start=None, deterministic=True):
|
||||
observation = torch.tensor(observations['observation'])
|
||||
safe_actions = torch.tensor(observations['safe_actions'])
|
||||
if deterministic:
|
||||
_, actions = self.forward(observation, safe_actions).max(-1)
|
||||
else:
|
||||
actions = self.sample(self.forward(observation, safe_actions))
|
||||
return actions, None
|
||||
|
||||
# def torch_dist_nomask(self, dist):
|
||||
# print('no mask logprob')
|
||||
# logits = dist[..., :self.action_dim]
|
||||
|
||||
@@ -1,6 +1,6 @@
|
||||
import torch
|
||||
from core.value_estimation import gae
|
||||
from core.optimization import conjugate_gradient, line_search
|
||||
from src.core.value_estimation import gae
|
||||
from src.core.optimization import conjugate_gradient, line_search
|
||||
|
||||
def trpo_step(value, policy, states, safe_actions, actions, rewards, dones, gamma, gae_lambda, delta, backtrack_coeff, backtrack_iters, v_opt, v_iters, cg_iters=10, cg_damping=0.1):
|
||||
|
||||
|
||||
@@ -9,6 +9,10 @@ class TransformObservation(gym.wrappers.TransformObservation):
|
||||
def __getattr__(self, name):
|
||||
return getattr(self.env, name)
|
||||
|
||||
class OptionsTimeLimit(gym.wrappers.TimeLimit):
|
||||
def __getattr__(self, name):
|
||||
return getattr(self.env, name)
|
||||
|
||||
class CollisionPenaltyWrapper(Wrapper):
|
||||
|
||||
def __init__(self, env, collision_distance, collision_penalty, *args, **kwargs):
|
||||
|
||||
Reference in New Issue
Block a user