Parameterize number of discriminator updates per epoch

This commit is contained in:
ebuehrle
2021-10-28 17:58:34 +02:00
parent 070b8fc785
commit b1740764e3

View File

@@ -9,9 +9,10 @@ def flatten_transitions(transitions):
'dones': np.stack(list(t['dones'] for t in transitions), axis=0), 'dones': np.stack(list(t['dones'] for t in transitions), axis=0),
} }
def train_discriminator(env, generator, discriminator, num_samples): def train_discriminator(env, generator, discriminator, num_samples, n_updates=1):
transitions = list(itertools.islice(env.sample_ll(generator), num_samples)) transitions = list(itertools.islice(env.sample_ll(generator), num_samples))
generator_samples = flatten_transitions(transitions) generator_samples = flatten_transitions(transitions)
for _ in range(n_updates):
discriminator.train_disc(gen_samples=generator_samples) discriminator.train_disc(gen_samples=generator_samples)
def train_generator(env, generator, discriminator, num_samples): def train_generator(env, generator, discriminator, num_samples):