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 self.steps = 0
def _after_step(self): 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()), state=torch.tensor(self.s).unsqueeze(0).to(self.discriminator.discrim_net.device()),
action=torch.tensor([[self.a]]).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 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): def _transitions(self):
yield { yield {
'obs': self.obs, 'obs': self.obs,
'action': self.ch, 'action': self.ch.cpu(),
'reward': self.r.detach(), 'reward': self.r,
'episode_start': self.episode_start, 'episode_start': self.episode_start,
'value': self.value.detach(), 'value': self.value.detach(),
'log_prob': self.log_prob.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]: for s in generator_samples[:-1]:
generator.rollout_buffer.add( generator.rollout_buffer.add(
obs=s['obs'], obs=s['obs'],
action=s['action'].cpu(), action=s['action'],
reward=s['reward'].cpu(), reward=s['reward'],
episode_start=s['episode_start'], episode_start=s['episode_start'],
value=s['value'], value=s['value'],
log_prob=s['log_prob'], log_prob=s['log_prob'],