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