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