Tensorflow로 50줄짜리 Original GAN Code 구현하기


GAN Implementation in 50 Lines of Tensorflow Code


코드는 이형민군의 깃허브 코드를 참조하였습니다. 맨 처음 GAN을 공부하실 때 도움이 될 것으로 희망합니다. Pytorch 코드는 여기를 참조하세요.

우선 Full-code는 맨 아래에서 정리하도록 하겠습니다.


1. GAN Easy Review


여기서는 이론적인 내용은 전혀 없으며 100% 구현 중심의 설명이다. GAN을 간단하게만 Review해 보면, 결국 다음 두 Network를 정의한다.

? Noise z를 받아서 Fake Data를 만드는 Generator

? Real Data와 Fake Data를 구분하는 Discriminator

결국 중요한 것은 Loss Function과 Training 방법일 것이다.

Loss Function을 정의하기 전에, 우선 Loss라는 것은 어떠한 System이 있을 때, 그것의 최종 결과물 (혹은 중간 결과물)을 기반으로 설정해주는 것이다.

여기서는 System의 최종 결과물은 Discriminator가 내뱉는 넌 진짜야, 넌 가짜야 하는 것인데.

? 진짜 = Real Data = 1

? 가짜 = Fake Data = 0

으로 정의한다. 0과 1 사이의 값이기 때문에 Cross Entropy로 Loss를 정의하면 되겠다라고 생각해야 한다. (머신러닝 기초를 공부한 사람이라면!)

그렇다면 본 System을 구성하는 Generator와 Discriminator의 Loss는 다음과 같이 된다.

? Discriminator는 Real Data와 1 사이의 Loss (= Cross Entropy)

? Discriminator는 Fake Data와 0 사이의 Loss (= Cross Entropy)

? Generator는 Fake Data와 1 사이의 Loss (= Cross Entropy, 속여야 하니까)

그렇다면 Cross Entropy 수식은?

CE(pred,y)=−ylog(pred)−(1−y)⋅log(1−pred)CE(pred,y)=−y⋅log(pred)−(1−y)⋅log(1−pred)

y는 실제 값, pred는 Network의 예측 값이다.

이것을 Discriminator에 적용하면 (그냥 대입이다)

DLoss=CE(D(x),1)+CE(D(G(z)),0)DLoss=CE(D(x),1)+CE(D(G(z)),0)=−log(D(x))−log(1−D(g(z))=−log(D(x))−log(1−D(g(z))

Generator에 적용하면

GLoss=CE(D(G(z)),1)GLoss=CE(D(G(z)),1)=−log(D(G(z))=−log(D(G(z))

끝이다. 이제 구현하자.


2. Implementation


50줄의 전체 코드는 크게 5개의 파트로 나누어진다.


2.1 Requirements and Dataset

본 코드는 Matplotlib과 Numpy가 필요하고, Tensorflow에서 제공하는 MNIST Dataset을 사용할 예정이다.

또한 입출력단에 사용되는 Placeholder로는 Image Input X (28x28)과 Noise Z (Dimension은 128을 가정)이 필요하다.

다음은 그 구현이다

import tensorflow as tf

import matplotlib.pyplot as plt

import numpy as np

from tensorflow.examples.tutorials.mnist import input_data


mnist = input_data.read_data_sets("./mnist/data/", one_hot=True)

X = tf.placeholder(tf.float32, [None, 28 * 28]) # MNIST = 28*28

Z = tf.placeholder(tf.float32, [None, 128]) # Noise Dimension = 128



2.2 Generator Network

Generator의 입력은 Noise Z이고. Fully-Connected 2-Layer를 지나 Real Image와 똑같은 28x28 Dimension의 출력을 내뱉는다.

Hidden Node 갯수는 256으로 정의한다.

MNIST는 0과 1사이의 흑백 Image이므로, 출력시 Sigmoid를 적용해 준다.

# ********* G-Network (Hidden Node # = 256)

G_W1 = tf.Variable(tf.random_normal([128, 256], stddev=0.01))

G_W2 = tf.Variable(tf.random_normal([256, 28 * 28], stddev=0.01))

G_b1 = tf.Variable(tf.zeros([256]))

G_b2 = tf.Variable(tf.zeros([28 * 28]))


defgenerator(noise_z):# 128 -> 256 -> 28*28

    hidden = tf.nn.relu(tf.matmul(noise_z, G_W1) + G_b1)

    output = tf.nn.sigmoid(tf.matmul(hidden, G_W2) + G_b2)

    return output



2.3 Discriminator Network

Generator의 반대로 보면 된다. 입력은 Image (28x28, 가짜든 진짜든)이고, 출력은 0과 1사이의 스칼라 값이다. 이것이 진짜와 가짜를 결정한다.

그러므로 출력단에 Sigmoid 넣어 주면 된다.

# ********* D-Network (Hidden Node # = 256)

D_W1 = tf.Variable(tf.random_normal([28 * 28, 256], stddev=0.01))

D_W2 = tf.Variable(tf.random_normal([256, 1], stddev=0.01))

D_b1 = tf.Variable(tf.zeros([256]))

D_b2 = tf.Variable(tf.zeros([1]))


defdiscriminator(inputs):# 28*28 -> 256 -> 1

    hidden = tf.nn.relu(tf.matmul(inputs, D_W1) + D_b1)

    output = tf.nn.sigmoid(tf.matmul(hidden, D_W2) + D_b2)

    return output



2.4 Generate Fake Image, Loss and Optimization

Noise를 Generator에 집어넣어 Fake Image를 만들고. 위에서 정의한 Loss를 그대로 구현하고.

Optimizer로는 Adam을 쓰자.

주의할 점은 Generator와 Discriminator는 서로를 건드리지 않고 별개로 Train된다는 것이다.

# ********* Generation, Loss, Optimization

G = generator(Z)


loss_D = -tf.reduce_mean(tf.log(discriminator(X)) + tf.log(1 - discriminator(G)))

loss_G = -tf.reduce_mean(tf.log(discriminator(G)))


train_D = tf.train.AdamOptimizer(learning_rate=0.0002).minimize(loss_D, var_list=[D_W1, D_b1, D_W2, D_b2])

train_G = tf.train.AdamOptimizer(learning_rate=0.0002).minimize(loss_G, var_list=[G_W1, G_b1, G_W2, G_b2])



2.5 Training and Testing

Tensorflow니까 Session을 열고 초기화를 해 준다.

Test시에는 고정된 Input (Noise)가 들어갔을 때 Generated되는 Image를 보고 싶기 때문에 Training Loop 밖에서 만들어 준다 (noise_test).

경축! 아무것도 안하여 에스천사게임즈가 새로운 모습으로 재오픈 하였습니다.
어린이용이며, 설치가 필요없는 브라우저 게임입니다.
https://s1004games.com

Epoch는 200번, Batch Size는 100으로 설정한다.

각 Batch마다 Noise를 만들고 Discriminator와 Generator를 번갈아가며 Training하면 Training은 끝이다.

Test과정에서는 noise_test를 Generator에 넣어주고 결과를 Plot하는 것이 전부이다.

sess = tf.Session()

sess.run(tf.global_variables_initializer())


# ********* Training and Testing

noise_test = np.random.normal(size=(10, 128)) # 10 = Test Sample Size, 128 = Noise Dimension

for epoch in range(200): # 200 = Num. of Epoch

    for i in range(int(mnist.train.num_examples / 100)): # 100 = Batch Size

        batch_xs, _ = mnist.train.next_batch(100)

        noise = np.random.normal(size=(100, 128))


        sess.run(train_D, feed_dict={X: batch_xs, Z: noise})

        sess.run(train_G, feed_dict={Z: noise})


    if epoch == 0or (epoch + 1) % 10 == 0: # 10 = Saving Period

        samples = sess.run(G, feed_dict={Z: noise_test})


        fig, ax = plt.subplots(1, 10, figsize=(10, 1))

        for i in range(10):

            ax[i].set_axis_off()

            ax[i].imshow(np.reshape(samples[i], (28, 28)))

        plt.savefig('samples_ex/{}.png'.format(str(epoch).zfill(3)), bbox_inches='tight')

        plt.close(fig)



3. Full Code


import tensorflow as tf

import matplotlib.pyplot as plt

import numpy as np

from tensorflow.examples.tutorials.mnist import input_data


mnist = input_data.read_data_sets("./mnist/data/", one_hot=True)

X = tf.placeholder(tf.float32, [None, 28 * 28]) # MNIST = 28*28

Z = tf.placeholder(tf.float32, [None, 128]) # Noise Dimension = 128


# ********* G-Network (Hidden Node # = 256)

G_W1 = tf.Variable(tf.random_normal([128, 256], stddev=0.01))

G_W2 = tf.Variable(tf.random_normal([256, 28 * 28], stddev=0.01))

G_b1 = tf.Variable(tf.zeros([256]))

G_b2 = tf.Variable(tf.zeros([28 * 28]))


defgenerator(noise_z):# 128 -> 256 -> 28*28

    hidden = tf.nn.relu(tf.matmul(noise_z, G_W1) + G_b1)

    output = tf.nn.sigmoid(tf.matmul(hidden, G_W2) + G_b2)

    return output


# ********* D-Network (Hidden Node # = 256)

D_W1 = tf.Variable(tf.random_normal([28 * 28, 256], stddev=0.01))

D_W2 = tf.Variable(tf.random_normal([256, 1], stddev=0.01))

D_b1 = tf.Variable(tf.zeros([256]))

D_b2 = tf.Variable(tf.zeros([1]))


defdiscriminator(inputs):# 28*28 -> 256 -> 1

    hidden = tf.nn.relu(tf.matmul(inputs, D_W1) + D_b1)

    output = tf.nn.sigmoid(tf.matmul(hidden, D_W2) + D_b2)

    return output


# ********* Generation, Loss, Optimization and Session Init.

G = generator(Z)

loss_D = -tf.reduce_mean(tf.log(discriminator(X)) + tf.log(1 - discriminator(G)))

loss_G = -tf.reduce_mean(tf.log(discriminator(G)))

train_D = tf.train.AdamOptimizer(learning_rate=0.0002).minimize(loss_D, var_list=[D_W1, D_b1, D_W2, D_b2])

train_G = tf.train.AdamOptimizer(learning_rate=0.0002).minimize(loss_G, var_list=[G_W1, G_b1, G_W2, G_b2])


sess = tf.Session()

sess.run(tf.global_variables_initializer())


# ********* Training and Testing

noise_test = np.random.normal(size=(10, 128)) # 10 = Test Sample Size, 128 = Noise Dimension

for epoch in range(200): # 200 = Num. of Epoch

    for i in range(int(mnist.train.num_examples / 100)): # 100 = Batch Size

        batch_xs, _ = mnist.train.next_batch(100)

        noise = np.random.normal(size=(100, 128))


        sess.run(train_D, feed_dict={X: batch_xs, Z: noise})

        sess.run(train_G, feed_dict={Z: noise})


    if epoch == 0or (epoch + 1) % 10 == 0: # 10 = Saving Period

        samples = sess.run(G, feed_dict={Z: noise_test})


        fig, ax = plt.subplots(1, 10, figsize=(10, 1))

        for i in range(10):

            ax[i].set_axis_off()

            ax[i].imshow(np.reshape(samples[i], (28, 28)))

        plt.savefig('samples_ex/{}.png'.format(str(epoch).zfill(3)), bbox_inches='tight')

        plt.close(fig)



결과는? 사실 썩 만족스러운 결과는 아닐 것이다. 하지만 이 코드는 Tutorial에 불과하고. GAN Original Paper에서 제시하는 Loss와 구현 방식을 그대로 적용한 예제로 보면 될 것 같다.

GAN의 수학적인 안정성에 관심이 많은 사람은 DCGAN을 거쳐 InfoGAN, f-GAN, EBGAN, WGAN, BEGAN, WGAN-GP 등을 보면 될 것 같다.

GAN의 응용 분야에 관심이 많은 사람은 DCGAN, VAE, InfoGAN을 공부한 뒤 Pix2Pix, CycleGAN, DiscoGAN을 보고 다음 포스팅을 참고하면 좋을 것 같다.


[출처] 

https://taeoh-kim.github.io/blog/tensorflow%EB%A1%9C-50%EC%A4%84%EC%A7%9C%EB%A6%AC-original-gan-code-%EA%B5%AC%ED%98%84%ED%95%98%EA%B8%B0/

본 웹사이트는 광고를 포함하고 있습니다.
광고 클릭에서 발생하는 수익금은 모두 웹사이트 서버의 유지 및 관리, 그리고 기술 콘텐츠 향상을 위해 쓰여집니다.
번호 제목 글쓴이 날짜 조회 수
공지 오라클 기본 샘플 데이터베이스 졸리운_곰 2014.01.02 86311
공지 [SQL컨셉] 서적 "SQL컨셉"의 샘플 데이타 베이스 SAMPLE DATABASE of ORACLE 가을의 곰을... 2013.02.10 78766
공지 [G_SQL] Sample Database 가을의 곰을... 2012.05.20 95513
46 [데이터분석 & 데이터 사이언스] 수많은 데이터 사이언티스트들이 직장을 떠나는 이유는 무엇인가? file 졸리운_곰 2025.03.09 901
45 [데이터분석][파이썬][python] 한글 글꼴 사용 (matplotlib) 졸리운_곰 2024.04.18 1360
44 [데이터분석 & 데이터 사이언스] 데이터에 관한 꼭 알아야 할 오해와 진실 12가지 졸리운_곰 2024.01.17 1315
43 [데이터분석][파이썬][python] Awesome Dash Awesome file 졸리운_곰 2021.07.10 2340
42 [데이터분석][파이썬][python] ???? Introducing Dash ???? file 졸리운_곰 2021.07.10 1607
41 [dataset] (한글) 욕설 감지 데이터셋 file 졸리운_곰 2021.05.12 1559
40 [데이터분석][python] Dash를 사용하는 초보자 및 기타 모든 사용자를위한 Python의 대시 보드 file 졸리운_곰 2021.04.14 1762
39 [데이터분석][python] Dash를 사용하는 초보자 및 기타 모든 사용자를위한 Python의 대시 보드 file 졸리운_곰 2021.04.14 1605
38 [데이터분석][데이터 사이언스][python][Dash] Python, Dash 및 Plotly를 사용하여 COVID-19 사례 데이터 시각화 file 졸리운_곰 2021.03.28 1501
37 [데이터분석][머신러닝] When not to use machine learning or AI Adventures in wishful thinking, nonstationarity, and pattern-finding / 기계 학습 또는 AI를 사용하지 않아야하는 경우 희망찬 사고, 비정상 성, 패턴 찾기의 모험 file 졸리운_곰 2021.03.28 21622
36 [MSA][머신러닝] 쿠버네티스 기반의 End2End 머신러닝 플랫폼 Kubeflow #1 - 소개 file 졸리운_곰 2021.03.21 1136
35 [데이터사이언스] 데이터 과학자를위한 3 가지 훌륭한 디자인 패턴, 3 Great Design Patterns for Data Scientists file 졸리운_곰 2021.03.04 780
34 [데이터분석] 시계열 데이터에 AI를 사용하는 이유는 무엇입니까? file 졸리운_곰 2021.02.28 1193
33 [데이터분석] AI 예측 및 이상 탐지를위한 시계열 데이터 전처리 file 졸리운_곰 2021.02.28 1031
32 [데이터분석] bitcoin analysis 비트 코인 시계열 데이터에 대한 AI 이상 탐지 file 졸리운_곰 2021.02.27 1592
31 [데이터분석 & 데이터 사이언스] How To Create a Data Science Portfolio Website file 졸리운_곰 2021.02.14 1824
30 [데이터수집4] 오픈 API 데이터 수집 (소셜미디어 데이터 수집) file 졸리운_곰 2020.06.12 1951
29 [데이터수집3] 관계형 데이터베이스 데이터 수집 file 졸리운_곰 2020.06.12 1318
28 [데이터수집2] 분산시스템 로그 수집 (빅데이터 수집) file 졸리운_곰 2020.06.12 1493
27 [데이터수집1] 웹 크롤링, 웹 스크래핑 file 졸리운_곰 2020.06.12 1804
대표 김성준 주소 : 경기 용인 분당수지 U타워 등록번호 : 142-07-27414
통신판매업 신고 : 제2012-용인수지-0185호 출판업 신고 : 수지구청 제 123호 개인정보보호최고책임자 : 김성준 sjkim70@stechstar.com
대표전화 : 010-4589-2193 [fax] 02-6280-1294 COPYRIGHT(C) stechstar.com ALL RIGHTS RESERVED