# -*- coding: utf-8 -*- """Pong_VAE Automatically generated by Colaboratory. Original file is located at https://colab.research.google.com/drive/1Q9Uhc5MgM8hkT_YpEQQw5OmpTWC2Z3Yd """ """What does a VAE of the observation space for Pong-v0 look like? Uses code from the autoencoder tut from keras """ from keras.datasets import mnist import gym from keras.layers import Lambda, Input, Dense, Conv2D, MaxPooling2D, UpSampling2D, Flatten, Reshape from keras.models import Model from keras.losses import mse, binary_crossentropy from keras.utils import plot_model from keras import backend as K from keras.callbacks import TensorBoard from IPython.display import SVG from keras.utils.vis_utils import model_to_dot import numpy as np import matplotlib.pyplot as plt import argparse import random import os # %matplotlib inline def get_data2D(data_dim=(105,80), environment='Pong-v0', number_of_examples=1000, render=False): """ Runs a bunch of random actions in the chosen Atari environment and screen captures the downscaled/grayscaled images :param data_dim: product of dims of the data that you wanna store e.g. 210*160=33600 :param environment: Defaults to Pong-v0 environment, adaptable to any Atari game in Gym :param number_of_examples: 1000 examples is pretty quick to generate :param render: Do you want to render this? :return: the data that we just made """ env = gym.make(environment) observation = env.reset() done = False episode_limit = 2000 data = np.zeros((number_of_examples, 8400), dtype=float) counter = 0 while counter < number_of_examples - 1: for t in range(episode_limit): if render: env.render() action = env.action_space.sample() observation, reward, done, info = env.step(action) buffer = observation.astype('float32') / 255 buffer = buffer[::2, ::2] buffer = rgb2_gray(buffer) buffer = np.reshape(buffer, [-1, 8400]) # number_of_examples is 1 indexed but counter is zero indexed... # a lot of the sequential observations are the same pixel values # checks to make sure it's new info if counter < (number_of_examples - 1): if np.sum(buffer - data[counter - 1]) != 0.0: data[counter] = buffer counter += 1 if done: print("Episode finished after {} timesteps".format(t + 1)) env.reset() print("Counter is {}".format(counter)) return data def rgb2_gray(rgb): r, g, b = rgb[:,:,0], rgb[:,:,1], rgb[:,:,2] gray = 0.2989 * r + 0.5870 * g + 0.1140 * b return gray x_train = get_data2D(number_of_examples=10000) x_test = get_data2D(number_of_examples=1000) plt.figure(figsize=(1,1)) n = 4 f, axarr = plt.subplots(1,n) for i in range(n): axarr[i].imshow(x_train[random.randint(0,len(x_train))].reshape(105,80)) axarr[i].axis('off') plt.show() def sampling(args): """reparameterization trick to sample from an isotropic unit gaussian :param args: (tensor) mean and log of variance Q(z|X) :return: z (tensor) sampled latent vector """ z_mean, z_log_var = args batch = K.shape(z_mean)[0] dim = K.int_shape(z_mean)[1] epsilon = K.random_normal(shape=(batch,dim)) return z_mean + K.exp(0.5 * z_log_var) * epsilon # network parameters original_dim = 8400 input_shape = x_train[0].shape intermediate_dim = 512 batch_size = 1024 latent_dim = 10 epochs = 50 # build encoder model inputs = Input(shape=input_shape, name='encoder_input') x= Dense(intermediate_dim, activation='relu', name='intermediate_encoder1')(inputs) #x= Dense(intermediate_dim, activation='relu', name='intermediate_encoder2')(x) z_mean = Dense(latent_dim, name='z_mean')(x) z_log_var = Dense(latent_dim, name='z_log_var')(x) z = Lambda(sampling, output_shape=(latent_dim,), name='z')([z_mean, z_log_var]) encoder = Model(inputs, [z_mean, z_log_var, z], name='encoder') encoder.summary() # build decoder model latent_inputs = Input(shape=(latent_dim,), name='z_sampling') x = Dense(intermediate_dim, activation='relu', name='intermediate_decoder1')(latent_inputs) #x = Dense(intermediate_dim, activation='relu', name='intermediate_decoder2')(x) outputs = Dense(original_dim, activation='sigmoid', name='outputs')(x) # instantiate decoder model decoder = Model(latent_inputs, outputs, name='decoder') decoder.summary() # instantiate VAE model outputs = decoder(encoder(inputs)[2]) vae = Model(inputs, outputs, name='VAE_MLP') beta = 1. reconstruction_loss = binary_crossentropy(inputs, outputs) reconstruction_loss *= original_dim kl_loss = 1 + z_log_var - K.square(z_mean) - K.exp(z_log_var) kl_loss = K.sum(kl_loss, axis=-1) kl_loss *= -0.5 vae_loss = K.mean(reconstruction_loss + beta * kl_loss) vae.add_loss(vae_loss) vae.compile(optimizer='adam') vae.summary() vae.fit(x_train, epochs=epochs, batch_size=batch_size, shuffle=True, validation_data=(x_test, None)) vae.save_weights('vae_mlp_pong.h5') z_mean, z_log_sigma, z = encoder.predict(x_test, batch_size=batch_size) decoded_img = decoder.predict(z_mean) plt.figure(figsize=(5,4)) n = 4 f, axarr = plt.subplots(2,n) for i in range(n): index = random.randint(0,len(x_test)) axarr[0,i].imshow(x_train[index].reshape(105,80)) axarr[0,i].axis('off') axarr[1,i].imshow(decoded_img[index].reshape(105,80)) axarr[1,i].axis('off') plt.show()