Integrate options env and policy

This commit is contained in:
ebuehrle
2022-02-15 18:36:53 +01:00
parent 072c0ff417
commit c6a4c10605
5 changed files with 60 additions and 9 deletions

View File

@@ -13,4 +13,4 @@ python -m src.eval_main
# idm # idm
python -m src.eval_main --method=idm python -m src.eval_main --method=idm
python -m src.eval_main --method=ogail --policy_file='checkpoints/gail-options-setobs2.pt' --env='NormalizedOptionsEvalEnv'

View File

@@ -15,8 +15,13 @@ class BasePolicy(nn.Module):
def sample(self, dist): def sample(self, dist):
return self.torch_dist(dist).sample() return self.torch_dist(dist).sample()
def predict(self, states): def predict(self, observations, state=None, episode_start=None, deterministic=True):
return self.sample(self.forward(states)) observations = torch.tensor(observations)
if deterministic:
actions = self.forward(observations)[..., :self.action_dim]
else:
actions = self.sample(self.forward(observations))
return actions, None
def log_prob(self, dist, actions): def log_prob(self, dist, actions):
return self.torch_dist(dist).log_prob(actions) return self.torch_dist(dist).log_prob(actions)
@@ -59,6 +64,14 @@ class DiscretePolicy(BasePolicy):
def torch_dist(self, dist): def torch_dist(self, dist):
return Categorical(logits=dist) return Categorical(logits=dist)
def predict(self, observations, state=None, episode_start=None, deterministic=True):
observations = torch.tensor(observations)
if deterministic:
_, actions = self.forward(observations).max(-1)
else:
actions = self.sample(self.forward(observations))
return actions, None
class SetPolicy(Policy): class SetPolicy(Policy):
def forward(self, states): def forward(self, states):

View File

@@ -8,6 +8,9 @@ from src.baselines import IDMRulePolicy
from src.evaluation import IntersimpleEvaluation from src.evaluation import IntersimpleEvaluation
import src.gail.options as options_envs import src.gail.options as options_envs
from src.evaluation.metrics import divergence, visualize_distribution from src.evaluation.metrics import divergence, visualize_distribution
from src.core.policy import SetPolicy, SetDiscretePolicy
from src.core.reparam_module import ReparamPolicy
from src.gail2 import envs as options_envs2
from typing import Optional, List, Dict, Tuple from typing import Optional, List, Dict, Tuple
import torch import torch
@@ -33,12 +36,31 @@ def load_policy(method:str,
if method == 'idm': if method == 'idm':
policy = IDMRulePolicy(env, **policy_kwargs) policy = IDMRulePolicy(env, **policy_kwargs)
elif method == 'bc': elif method == 'bc':
raise NotImplementedError policy = SetPolicy(env.action_space.shape[-1])
policy.load_state_dict(torch.load(policy_file))
policy.eval()
elif method == 'gail': elif method == 'gail':
policy = sb3.PPO.load(policy_file) policy = SetPolicy(env.action_space.shape[-1])
raise NotImplementedError policy(torch.zeros(env.observation_space.shape))
policy = ReparamPolicy(policy)
policy.load_state_dict(torch.load(policy_file))
policy.eval()
elif method == 'gail-ppo':
policy = SetPolicy(env.action_space.shape[-1])
policy.load_state_dict(torch.load(policy_file))
policy.eval()
elif method == 'rail': elif method == 'rail':
raise NotImplementedError raise NotImplementedError
elif method == 'ogail':
policy = SetDiscretePolicy(env.action_space.n)
policy(torch.zeros(env.observation_space.shape))
policy = ReparamPolicy(policy)
policy.load_state_dict(torch.load(policy_file))
policy.eval()
elif method == 'ogail-ppo':
policy = SetDiscretePolicy(env.action_space.n)
policy.load_state_dict(torch.load(policy_file))
policy.eval()
elif method == 'sgail': elif method == 'sgail':
policy = sb3.PPO.load(policy_file) policy = sb3.PPO.load(policy_file)
raise NotImplementedError raise NotImplementedError
@@ -158,6 +180,7 @@ def evaluate_policy(locations:List[Tuple[int,int]],
""" """
envs_dict = dict(intersim.envs.intersimple.__dict__) envs_dict = dict(intersim.envs.intersimple.__dict__)
envs_dict.update(dict(options_envs.__dict__)) envs_dict.update(dict(options_envs.__dict__))
envs_dict.update(dict(options_envs2.__dict__))
policy_metrics = [None]* len(locations) policy_metrics = [None]* len(locations)
# iterate through vehicles # iterate through vehicles

View File

@@ -6,6 +6,7 @@ from typing import Callable, Dict, Optional
import os import os
import pickle import pickle
from tqdm import tqdm from tqdm import tqdm
from src.gail2.envs import OptionsEnv
class IntersimpleEvaluation: class IntersimpleEvaluation:
""" """
@@ -34,6 +35,7 @@ class IntersimpleEvaluation:
self.env = eval_env self.env = eval_env
self.n_episodes = eval_env.nv self.n_episodes = eval_env.nv
self.use_pbar = use_pbar self.use_pbar = use_pbar
self.is_options_env = isinstance(self.env, OptionsEnv)
# metrics present on every step of every episode # metrics present on every step of every episode
self.metric_keys_all = ['v_all', 'a_all', 'col_all'] self.metric_keys_all = ['v_all', 'a_all', 'col_all']
@@ -85,11 +87,14 @@ class IntersimpleEvaluation:
if self.use_pbar: if self.use_pbar:
self.pbar = tqdm(total=self.n_episodes) self.pbar = tqdm(total=self.n_episodes)
if self.is_options_env:
print('Evaluating an options environment')
evaluate_policy( evaluate_policy(
policy, policy,
self.env, self.env,
n_eval_episodes=self.n_episodes, n_eval_episodes=self.n_episodes,
callback=self.evaluate_policy_callback, callback=self.evaluate_options_policy_callback if self.is_options_env else self.evaluate_policy_callback,
return_episode_rewards=False return_episode_rewards=False
) )
if self.use_pbar: if self.use_pbar:
@@ -100,6 +105,13 @@ class IntersimpleEvaluation:
self.save(filestr) self.save(filestr)
return self._metrics return self._metrics
def evaluate_options_policy_callback(self, local_vars, global_vars):
infos = local_vars['info']['ll']['infos']
dones = local_vars['info']['ll']['env_done']
agents = [info['agent'] for info in infos]
for info, done, agent in zip(infos, dones, agents):
self.eval_policy_step(info, done, agent)
def evaluate_policy_callback(self, local_vars, global_vars): def evaluate_policy_callback(self, local_vars, global_vars):
""" """
Callback run in evaluate_policy after taking an action and receiving an observation Callback run in evaluate_policy after taking an action and receiving an observation
@@ -112,6 +124,9 @@ class IntersimpleEvaluation:
env = local_vars['env'].envs[venv_i] env = local_vars['env'].envs[venv_i]
assert isinstance(env, Intersimple) assert isinstance(env, Intersimple)
self.eval_policy_step(info, done, _agent)
def eval_policy_step(self, info, done, _agent):
# Increase collision counter if episode terminated with a collision # Increase collision counter if episode terminated with a collision
self._metrics['v_all'][_agent].append(info['prev_state'][_agent,2].item()) self._metrics['v_all'][_agent].append(info['prev_state'][_agent,2].item())
self._metrics['a_all'][_agent].append(info['action_taken'][_agent,0].item()) self._metrics['a_all'][_agent].append(info['action_taken'][_agent,0].item())

View File

@@ -1,6 +1,6 @@
import gym import gym
import numpy as np import numpy as np
from wrappers import Setobs, TransformObservation from src.gail2.wrappers import Wrapper, Setobs, TransformObservation
from intersim.envs import IntersimpleLidarFlatIncrementingAgent from intersim.envs import IntersimpleLidarFlatIncrementingAgent
obs_min = np.array([ obs_min = np.array([
@@ -30,7 +30,7 @@ def NormalizedOptionsEvalEnv(**kwargs):
), lambda obs: (obs - obs_min) / (obs_max - obs_min + 1e-10)) ), lambda obs: (obs - obs_min) / (obs_max - obs_min + 1e-10))
), options=[(0, 5), (1, 5), (2, 5), (4, 5), (6, 5), (8, 5)]) ), options=[(0, 5), (1, 5), (2, 5), (4, 5), (6, 5), (8, 5)])
class OptionsEnv(gym.Wrapper): class OptionsEnv(Wrapper):
def __init__(self, env, options): def __init__(self, env, options):
super().__init__(env) super().__init__(env)