From 4928458e08e9d41b5d2f737eecf30c78c8e5efc4 Mon Sep 17 00:00:00 2001 From: ebuehrle <43623224+ebuehrle@users.noreply.github.com> Date: Wed, 15 Sep 2021 07:15:43 +0200 Subject: [PATCH] Fix discriminator reward Had wrong sign. --- scratch/etienne/intersimple/gail/discriminator.py | 6 ++++-- scratch/etienne/intersimple/gail_options_image.py | 8 +++++--- 2 files changed, 9 insertions(+), 5 deletions(-) diff --git a/scratch/etienne/intersimple/gail/discriminator.py b/scratch/etienne/intersimple/gail/discriminator.py index 3206dad..f3c16e0 100644 --- a/scratch/etienne/intersimple/gail/discriminator.py +++ b/scratch/etienne/intersimple/gail/discriminator.py @@ -36,7 +36,8 @@ class CnnDiscriminator(torch.nn.Module): def forward(self, state, action): sa = self._concatenate(state, action) - return self.cnn(sa).squeeze() + assert sa.ndim == 4 + return self.cnn(sa).squeeze(1) class MlpDiscriminator(torch.nn.Module): """MLP similar to stable_baselines3.common.policies.ActorCriticPolicy.""" @@ -55,4 +56,5 @@ class MlpDiscriminator(torch.nn.Module): def forward(self, state, action): flat = self.flatten(state) sa = torch.cat((action, flat), -1) - return self.mlp(sa).squeeze() + assert sa.ndim == 2 + return self.mlp(sa).squeeze(1) diff --git a/scratch/etienne/intersimple/gail_options_image.py b/scratch/etienne/intersimple/gail_options_image.py index 636cef3..ceddd17 100644 --- a/scratch/etienne/intersimple/gail_options_image.py +++ b/scratch/etienne/intersimple/gail_options_image.py @@ -135,9 +135,11 @@ class HLOptions(OptionsEnv): self.steps = 0 def _after_step(self): - self.r += self.discount**self.steps * self.discriminator.discrim_net.discriminator( - torch.tensor(self.s).unsqueeze(0).to(self.discriminator.discrim_net.device()), - torch.tensor([[self.a]]).to(self.discriminator.discrim_net.device()), + self.r += self.discount**self.steps * self.discriminator.discrim_net.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 + done=torch.tensor(self.done).unsqueeze(0).to(self.discriminator.discrim_net.device()), # unused ) self.steps += 1