# -*- 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()
Comments