From 8dd42abbf375d5d77818f9f5ac0944a3a9b9f8f9 Mon Sep 17 00:00:00 2001 From: ebuehrle <43623224+ebuehrle@users.noreply.github.com> Date: Fri, 29 Oct 2021 14:08:32 +0200 Subject: [PATCH] use discriminator preprocessing --- src/gail/options.py | 6 +++--- src/gail/train.py | 4 ++-- 2 files changed, 5 insertions(+), 5 deletions(-) diff --git a/src/gail/options.py b/src/gail/options.py index b39fe1f..b160852 100644 --- a/src/gail/options.py +++ b/src/gail/options.py @@ -101,7 +101,7 @@ class HLOptions(OptionsEnv): self.steps = 0 def _after_step(self): - self.r += self.discount**self.steps * self.discriminator.discrim_net.reward_train( + self.r += self.discount**self.steps * self.discriminator.discrim_net.predict_reward_train( state=torch.tensor(self.s).unsqueeze(0).to(self.discriminator.discrim_net.device()), action=torch.tensor([[self.a]]).to(self.discriminator.discrim_net.device()), next_state=torch.tensor(self.s).unsqueeze(0).to(self.discriminator.discrim_net.device()), # unused @@ -112,8 +112,8 @@ class HLOptions(OptionsEnv): def _transitions(self): yield { 'obs': self.obs, - 'action': self.ch, - 'reward': self.r.detach(), + 'action': self.ch.cpu(), + 'reward': self.r, 'episode_start': self.episode_start, 'value': self.value.detach(), 'log_prob': self.log_prob.detach(), diff --git a/src/gail/train.py b/src/gail/train.py index ee2715c..3937292 100644 --- a/src/gail/train.py +++ b/src/gail/train.py @@ -22,8 +22,8 @@ def train_generator(env, generator, discriminator, num_samples): for s in generator_samples[:-1]: generator.rollout_buffer.add( obs=s['obs'], - action=s['action'].cpu(), - reward=s['reward'].cpu(), + action=s['action'], + reward=s['reward'], episode_start=s['episode_start'], value=s['value'], log_prob=s['log_prob'],