Fix rendering
- support predict() - move discount to OptionsEnv - fix RenderOptions
This commit is contained in:
@@ -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
|
||||||
|
|
||||||
|
|||||||
@@ -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():
|
||||||
|
|||||||
@@ -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):
|
||||||
"""
|
"""
|
||||||
|
|||||||
Reference in New Issue
Block a user