use discriminator preprocessing

This commit is contained in:
ebuehrle
2021-10-29 14:08:32 +02:00
parent 92981ba284
commit 8dd42abbf3
2 changed files with 5 additions and 5 deletions

View File

@@ -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(),

View File

@@ -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'],