Fix rendering

- support predict()
- move discount to OptionsEnv
- fix RenderOptions
This commit is contained in:
ebuehrle
2021-11-02 16:15:34 +01:00
parent 6416fceb60
commit 2aaaad36f0
3 changed files with 30 additions and 19 deletions

View File

@@ -15,7 +15,7 @@ def imitation_discriminator(discriminator):
class OptionsEnv(gym.Wrapper): class OptionsEnv(gym.Wrapper):
def __init__(self, env, options, discriminator, ll_buffer_capacity, *args, **kwargs): def __init__(self, env, options, discriminator, discount, ll_buffer_capacity, *args, **kwargs):
super().__init__(env, *args, **kwargs) super().__init__(env, *args, **kwargs)
self.options = options self.options = options
@@ -27,6 +27,7 @@ class OptionsEnv(gym.Wrapper):
}) })
self.discriminator = discriminator self.discriminator = discriminator
self.discount = discount
self.ll_buffer_capacity = ll_buffer_capacity self.ll_buffer_capacity = ll_buffer_capacity
self.ll_buffer = deque(maxlen=ll_buffer_capacity) self.ll_buffer = deque(maxlen=ll_buffer_capacity)
@@ -85,11 +86,11 @@ class OptionsEnv(gym.Wrapper):
class RenderOptions(OptionsEnv): class RenderOptions(OptionsEnv):
def __init__(self, options, *args, **kwargs): def __init__(self, env, options, *args, **kwargs):
super().__init__(options, discriminator=lambda s, a, n, d: 0, ll_buffer_capacity=0, *args, **kwargs) super().__init__(env, options, discriminator=lambda s, a, n, d: 0, discount=1, ll_buffer_capacity=0, *args, **kwargs)
def _ll_step(self): def _ll_step(self, action):
out = super()._ll_step() out = super()._ll_step(action)
self.env.render() self.env.render()
return out return out

View File

@@ -34,7 +34,6 @@ def train(
epochs=200, epochs=200,
): ):
env = NRasterizedRouteSpeedRandomAgentLocation(**env_settings) env = NRasterizedRouteSpeedRandomAgentLocation(**env_settings)
env.discount = discount
tempdir = tempfile.TemporaryDirectory(prefix="quickstart") tempdir = tempfile.TemporaryDirectory(prefix="quickstart")
tempdir_path = pathlib.Path(tempdir.name) tempdir_path = pathlib.Path(tempdir.name)
@@ -51,10 +50,12 @@ def train(
gen_algo=stable_baselines3.PPO("CnnPolicy", venv), # unused gen_algo=stable_baselines3.PPO("CnnPolicy", venv), # unused
) )
options_env = TimeLimit(OptionsEnv(env, options_env = TimeLimit(OptionsEnv(
discriminator=imitation_discriminator(discriminator), env,
options=ALL_OPTIONS, options=ALL_OPTIONS,
ll_buffer_capacity=expert_batch_size discriminator=imitation_discriminator(discriminator),
discount=discount,
ll_buffer_capacity=expert_batch_size,
), max_episode_steps=10) ), max_episode_steps=10)
generator = stable_baselines3.PPO( generator = stable_baselines3.PPO(
OptionsCnnPolicy, OptionsCnnPolicy,
@@ -79,11 +80,15 @@ def train(
return generator return generator
def video(model_name, env): def video(model_name, env):
model = stable_baselines3.PPO.load(model_name)
env = RenderOptions(env, options=ALL_OPTIONS) env = RenderOptions(env, options=ALL_OPTIONS)
for s in env.sample_ll(model): model = stable_baselines3.PPO.load(model_name)
if s['dones']:
break done = False
obs = env.reset()
while not done:
action, _ = model.predict(obs)
obs, _, done, _ = env.step(action)
env.close(filestr='render/'+model_name) env.close(filestr='render/'+model_name)
def evaluate(): def evaluate():

View File

@@ -1,12 +1,13 @@
import stable_baselines3 from stable_baselines3.common.policies import ActorCriticPolicy, ActorCriticCnnPolicy
from torch.distributions import Categorical from torch.distributions import Categorical
class OptionsCnnPolicy(stable_baselines3.common.policies.ActorCriticCnnPolicy): class OptionsCnnPolicy(ActorCriticPolicy):
""" """
Class for high-level options policy (generator) Class for high-level options policy (generator)
""" """
def __init__(self, observation_space, *args, **kwargs): def __init__(self, observation_space, *args, **kwargs):
super().__init__(observation_space['obs'], *args, **kwargs) super().__init__(observation_space, *args, **kwargs)
self.cnn_policy = ActorCriticCnnPolicy(observation_space['obs'], *args, **kwargs)
def _prior_distribution(self, s): def _prior_distribution(self, s):
""" """
@@ -17,9 +18,9 @@ class OptionsCnnPolicy(stable_baselines3.common.policies.ActorCriticCnnPolicy):
values (torch.tensor): values from critic values (torch.tensor): values from critic
dist (torch.distributions): prior distribution over actions dist (torch.distributions): prior distribution over actions
""" """
latent_pi, latent_vf, latent_sde = self._get_latent(s) latent_pi, latent_vf, latent_sde = self.cnn_policy._get_latent(s)
distribution = self._get_action_dist_from_latent(latent_pi, latent_sde) distribution = self.cnn_policy._get_action_dist_from_latent(latent_pi, latent_sde)
values = self.value_net(latent_vf) values = self.cnn_policy.value_net(latent_vf)
return values, distribution.distribution return values, distribution.distribution
def forward(self, obs, eps=1e-6): def forward(self, obs, eps=1e-6):
@@ -39,6 +40,10 @@ class OptionsCnnPolicy(stable_baselines3.common.policies.ActorCriticCnnPolicy):
posterior = Categorical((prior.probs + eps) * m) posterior = Categorical((prior.probs + eps) * m)
ch = posterior.sample() ch = posterior.sample()
return ch, values, posterior.log_prob(ch) return ch, values, posterior.log_prob(ch)
def _predict(self, obs, deterministic=False):
action, _, _ = self.forward(obs)
return action
def evaluate_actions(self, obs, ch, eps=1e-6): def evaluate_actions(self, obs, ch, eps=1e-6):
""" """