woodsja icon

pong_VAE.py

woodsja | PRO | 01/15/19 12:52:02 AM UTC | 0 ⭐ | 325 👁️ | Never ⏰ | []
Python |

5.62 KB

|

None

|

0 👍

/

0 👎

# -*- 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