37 lines
1.3 KiB
Python
37 lines
1.3 KiB
Python
import numpy as np
|
|
import itertools
|
|
|
|
def flatten_transitions(transitions):
|
|
return {
|
|
'obs': np.stack(list(t['obs'] for t in transitions), axis=0),
|
|
'next_obs': np.stack(list(t['next_obs'] for t in transitions), axis=0),
|
|
'acts': np.stack(list(t['acts'] 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):
|
|
transitions = list(itertools.islice(env.sample_ll(generator), num_samples))
|
|
generator_samples = flatten_transitions(transitions)
|
|
discriminator.train_disc(gen_samples=generator_samples)
|
|
|
|
def train_generator(env, generator, discriminator, num_samples):
|
|
generator_samples = list(itertools.islice(env.sample_hl(generator, discriminator), num_samples+1))
|
|
|
|
generator.rollout_buffer.reset()
|
|
for s in generator_samples[:-1]:
|
|
generator.rollout_buffer.add(
|
|
obs=s['obs'],
|
|
action=s['action'].cpu(),
|
|
reward=s['reward'].cpu(),
|
|
episode_start=s['episode_start'],
|
|
value=s['value'],
|
|
log_prob=s['log_prob'],
|
|
)
|
|
|
|
generator.rollout_buffer.compute_returns_and_advantage(
|
|
last_values=generator_samples[-1]['value'],
|
|
dones=generator_samples[-1]['done'],
|
|
)
|
|
|
|
generator.train()
|