Switch back to additive reward

This commit is contained in:
ebuehrle
2021-09-13 17:55:50 +02:00
parent b0b358544f
commit 7eae74a7d8

View File

@@ -128,15 +128,15 @@ def sample(env, generator, discriminator, level: str):
assert feasible(env, plan, ch), f'Infeasible hl action {ch}' assert feasible(env, plan, ch), f'Infeasible hl action {ch}'
r = 0 r = 0
steps = 0 discount = 1
while not done and plan and feasible(env, plan, ch): while not done and plan and feasible(env, plan, ch):
a, plan = env._normalize(plan[0]), plan[1:] a, plan = env._normalize(plan[0]), plan[1:]
if level == 'high': if level == 'high':
r += discriminator.discrim_net.discriminator( r += discount * discriminator.discrim_net.discriminator(
torch.tensor(s).unsqueeze(0).to(discriminator.discrim_net.device()), torch.tensor(s).unsqueeze(0).to(discriminator.discrim_net.device()),
torch.tensor([[a]]).to(discriminator.discrim_net.device()), torch.tensor([[a]]).to(discriminator.discrim_net.device()),
) )
steps += 1 discount *= env.discount
nexts, _, done, _ = env.step(a) nexts, _, done, _ = env.step(a)
m = available_actions(env) m = available_actions(env)
@@ -151,11 +151,10 @@ def sample(env, generator, discriminator, level: str):
s = nexts s = nexts
if level == 'high': if level == 'high':
assert steps > 0
yield { yield {
'obs': obs, 'obs': obs,
'option': ch, 'option': ch,
'reward': r.detach() / steps, 'reward': r.detach(),
'episode_start': episode_start, 'episode_start': episode_start,
'value': value.detach(), 'value': value.detach(),
'log_prob': log_prob.detach(), 'log_prob': log_prob.detach(),
@@ -207,8 +206,9 @@ class OptionsEnv(gym.Wrapper):
'mask': gym.spaces.Box(low=0, high=1, shape=(num_hl_options,)), 'mask': gym.spaces.Box(low=0, high=1, shape=(num_hl_options,)),
}) })
def train(expert_data, epochs=10, expert_batch_size=32, generator_steps=2048): def train(expert_data, epochs=10, expert_batch_size=32, generator_steps=2048, discount=0.99):
env = NRasterized(**env_settings) env = NRasterized(**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)