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):
|
||||
|
||||
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)
|
||||
|
||||
self.options = options
|
||||
@@ -27,6 +27,7 @@ class OptionsEnv(gym.Wrapper):
|
||||
})
|
||||
|
||||
self.discriminator = discriminator
|
||||
self.discount = discount
|
||||
self.ll_buffer_capacity = ll_buffer_capacity
|
||||
self.ll_buffer = deque(maxlen=ll_buffer_capacity)
|
||||
|
||||
@@ -85,11 +86,11 @@ class OptionsEnv(gym.Wrapper):
|
||||
|
||||
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 __init__(self, env, options, *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):
|
||||
out = super()._ll_step()
|
||||
def _ll_step(self, action):
|
||||
out = super()._ll_step(action)
|
||||
self.env.render()
|
||||
return out
|
||||
|
||||
|
||||
@@ -34,7 +34,6 @@ def train(
|
||||
epochs=200,
|
||||
):
|
||||
env = NRasterizedRouteSpeedRandomAgentLocation(**env_settings)
|
||||
env.discount = discount
|
||||
|
||||
tempdir = tempfile.TemporaryDirectory(prefix="quickstart")
|
||||
tempdir_path = pathlib.Path(tempdir.name)
|
||||
@@ -51,10 +50,12 @@ def train(
|
||||
gen_algo=stable_baselines3.PPO("CnnPolicy", venv), # unused
|
||||
)
|
||||
|
||||
options_env = TimeLimit(OptionsEnv(env,
|
||||
discriminator=imitation_discriminator(discriminator),
|
||||
options_env = TimeLimit(OptionsEnv(
|
||||
env,
|
||||
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)
|
||||
generator = stable_baselines3.PPO(
|
||||
OptionsCnnPolicy,
|
||||
@@ -79,11 +80,15 @@ def train(
|
||||
return generator
|
||||
|
||||
def video(model_name, env):
|
||||
model = stable_baselines3.PPO.load(model_name)
|
||||
env = RenderOptions(env, options=ALL_OPTIONS)
|
||||
for s in env.sample_ll(model):
|
||||
if s['dones']:
|
||||
break
|
||||
model = stable_baselines3.PPO.load(model_name)
|
||||
|
||||
done = False
|
||||
obs = env.reset()
|
||||
while not done:
|
||||
action, _ = model.predict(obs)
|
||||
obs, _, done, _ = env.step(action)
|
||||
|
||||
env.close(filestr='render/'+model_name)
|
||||
|
||||
def evaluate():
|
||||
|
||||
Reference in New Issue
Block a user