텐서플로우(TensorFlow)를 이용해서 MNIST 데이터를 생성하는 GAN(Generative Adversarial Networks) 생성 모델(Generative Model) 구현해보기 – GAN 예제

이번 시간에는 텐서플로우(TensorFlow)를 이용해서 MNIST 데이터를 생성해내는 GAN(Generative Adversarial Networks) 모델을 구현해보자.

 

Generative Adversarial Networks(GAN)

Generative Adversarial Networks(GAN)은 적대적 학습(Adversarial Networks) 구조를 이용해서 생성 모델을 학습하는 아키텍쳐이다. GAN은 Discriminator(구분자)와 Generator(생성자)로 구성되어있다.

이는 경찰과 위조지폐생성범의 관계로 비유할 수 있다. 구분자는 어떤 이미지가 진짜 이미지인지 아니면 생성자가 만들어낸 가짜 이미지인지를 구분하도록 학습한다. 생성자는 Latent Variable(Noise Distribution에서 추출한 값)로부터 생성한 이미지가 구분자를 잘 속일 수 있도록 학습한다.

결과적으로 생성자는 원래 데이터의 분포(Distribution)을 거의 정확하게 근사하게 되고, 구분자는 50% 확률로 진짜 이미지와 생성자에 의해 생성된 가짜 이미지를 구분하는 균형점에 도달하게 된다. 새로운 이미지를 생성하고 싶을 때는 학습된 생성자(Generator)를 사용한다.

그림 1– Generative Adversarial Networks(GAN) architecture

그림 2 – GAN의 학습 과정(원래 데이터의 분포를 근사한다.)

 

텐서플로우(TensorFlow)를 이용한 MNIST 데이터 생성

이제 TensorFlow를 이용해서 MNIST 데이터의 분포를 학습하는 GAN 모델을 구현해보자.

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

  # -*- coding: utf-8 -*-
  """
  GAN(Generative Adversarial Networks)을 이용한 MNIST 데이터 생성
   
  Reference : https://github.com/TengdaHan/GAN-TensorFlow
   
  Author : solaris33
  Project URL : http://solarisailab.com/archives/2482
  """
   
  import tensorflow as tf
  import numpy as np
  import matplotlib.pyplot as plt
  import matplotlib.gridspec as gridspec
  import os
   
  # MNIST 데이터를 불러옵니다.
  from tensorflow.examples.tutorials.mnist import input_data
  mnist = input_data.read_data_sets("./mnist/data/", one_hot=True)
   
  # 생성된 MNIST 이미지를 8x8 Grid로 보여주는 plot 함수를 정의합니다.
  def plot(samples):
  fig = plt.figure(figsize=(8, 8))
  gs = gridspec.GridSpec(8, 8)
  gs.update(wspace=0.05, hspace=0.05)
   
  for i, sample in enumerate(samples):
  ax = plt.subplot(gs[i])
  plt.axis('off')
  plt.imshow(sample.reshape(28, 28))
  return fig
   
  # 설정값들을 선언합니다.
  num_epoch = 100000
  batch_size = 64
  num_input = 28 * 28
  num_latent_variable = 100
  num_hidden = 128
  learning_rate = 0.001
   
  # 플레이스 홀더를 선언합니다.
  X = tf.placeholder(tf.float32, [None, num_input]) # 인풋 이미지
  z = tf.placeholder(tf.float32, [None, num_latent_variable]) # 인풋 Latent Variable
   
  # Generator 변수들 설정
  # 100 -> 128 -> 784
  with tf.variable_scope('generator'):
  # 히든 레이어 파라미터
  G_W1 = tf.Variable(tf.random_normal(shape=[num_latent_variable, num_hidden], stddev=5e-2))
  G_b1 = tf.Variable(tf.constant(0.1, shape=[num_hidden]))
  # 아웃풋 레이어 파라미터
  G_W2 = tf.Variable(tf.random_normal(shape=[num_hidden, num_input], stddev=5e-2))
  G_b2 = tf.Variable(tf.constant(0.1, shape=[num_input]))
   
  # Discriminator 변수들 설정
  # 784 -> 128 -> 1
  with tf.variable_scope('discriminator'):
  # 히든 레이어 파라미터
  D_W1 = tf.Variable(tf.random_normal(shape=[num_input, num_hidden], stddev=5e-2))
  D_b1 = tf.Variable(tf.constant(0.1, shape=[num_hidden]))
  # 아웃풋 레이어 파라미터
  D_W2 = tf.Variable(tf.random_normal(shape=[num_hidden, 1], stddev=5e-2))
  D_b2 = tf.Variable(tf.constant(0.1, shape=[1]))
   
  # Generator를 생성하는 함수를 정의합니다.
  # Inputs:
  # X : 인풋 Latent Variable
  # Output:
  # generated_mnist_image : 생성된 MNIST 이미지
  def build_generator(X):
  hidden_layer = tf.nn.relu((tf.matmul(X, G_W1) + G_b1))
  output_layer = tf.matmul(hidden_layer, G_W2) + G_b2
  generated_mnist_image = tf.nn.sigmoid(output_layer)
   
  return generated_mnist_image
   
  # Discriminator를 생성하는 함수를 정의합니다.
  # Inputs:
  # X : 인풋 이미지
  # Output:
  # predicted_value : Discriminator가 판단한 True(1) or Fake(0)
  # logits : sigmoid를 씌우기전의 출력값
  def build_discriminator(X):
  hidden_layer = tf.nn.relu((tf.matmul(X, D_W1) + D_b1))
  logits = tf.matmul(hidden_layer, D_W2) + D_b2
  predicted_value = tf.nn.sigmoid(logits)
   
  return predicted_value, logits
   
  # 생성자(Generator)를 선언합니다.
  G = build_generator(z)
   
  # 구분자(Discriminator)를 선언합니다.
  D_real, D_real_logits = build_discriminator(X) # D(x)
  D_fake, D_fake_logits = build_discriminator(G) # D(G(z))
   
  # Discriminator의 손실 함수를 정의합니다.
  d_loss_real = tf.reduce_mean(tf.nn.sigmoid_cross_entropy_with_logits(logits=D_real_logits, labels=tf.ones_like(D_real_logits))) # log(D(x))
  d_loss_fake = tf.reduce_mean(tf.nn.sigmoid_cross_entropy_with_logits(logits=D_fake_logits, labels=tf.zeros_like(D_fake_logits))) # log(1-D(G(z)))
  d_loss = d_loss_real + d_loss_fake # log(D(x)) + log(1-D(G(z)))
   
  # Generator의 손실 함수를 정의합니다.
  g_loss = tf.reduce_mean(tf.nn.sigmoid_cross_entropy_with_logits(logits=D_fake_logits, labels=tf.ones_like(D_fake_logits))) # log(D(G(z))
   
  # 전체 파라미터를 Discriminator와 관련된 파라미터와 Generator와 관련된 파라미터로 나눕니다.
  tvar = tf.trainable_variables()
  dvar = [var for var in tvar if 'discriminator' in var.name]
  gvar = [var for var in tvar if 'generator' in var.name]
   
  # Discriminator와 Generator의 Optimizer를 정의합니다.
  d_train_step = tf.train.AdamOptimizer(learning_rate).minimize(d_loss, var_list=dvar)
  g_train_step = tf.train.AdamOptimizer(learning_rate).minimize(g_loss, var_list=gvar)
   
  # 생성된 이미지들을 저장할 generated_outputs 폴더를 생성합니다.
  num_img = 0
  if not os.path.exists('generated_output/'):
  os.makedirs('generated_output/')
   
  with tf.Session() as sess:
  # 변수들에 초기값을 할당합니다.
  sess.run(tf.global_variables_initializer())
   
  # num_epoch 횟수만큼 최적화를 수행합니다.
  for i in range(num_epoch):
  # MNIST 이미지를 batch_size만큼 불러옵니다.
  batch_X, _ = mnist.train.next_batch(batch_size)
  # Latent Variable의 인풋으로 사용할 noise를 Uniform Distribution에서 batch_size만큼 샘플링합니다.
  batch_noise = np.random.uniform(-1., 1., [batch_size, 100])
   
  # 500번 반복할때마다 생성된 이미지를 저장합니다.
  if i % 500 == 0:
  samples = sess.run(G, feed_dict={z: np.random.uniform(-1., 1., [64, 100])})
  fig = plot(samples)
  plt.savefig('generated_output/%s.png' % str(num_img).zfill(3), bbox_inches='tight')
  num_img += 1
  plt.close(fig)
   
  # Discriminator 최적화를 수행하고 Discriminator의 손실함수를 return합니다.
  _, d_loss_print = sess.run([d_train_step, d_loss], feed_dict={X: batch_X, z: batch_noise})
   
  # Generator 최적화를 수행하고 Generator 손실함수를 return합니다.
  _, g_loss_print = sess.run([g_train_step, g_loss], feed_dict={z: batch_noise})
   
  # 100번 반복할때마다 Discriminator의 손실함수와 Generator 손실함수를 출력합니다.
  if i % 100 == 0:
  print('반복(Epoch): %d, Generator 손실함수(g_loss): %f, Discriminator 손실함수(d_loss): %f' % (i, g_loss_print, d_loss_print))
view rawmnist_gan.py hosted with ❤ by GitHub

 

학습이 끝나면 아래와 같이 GAN 모델이 그럴듯한 MNIST 데이터를 생성해낸 모습을 볼 수 있다.

그림 3 – GAN이 생성해낸 MNIST 이미지

 

References

[1] https://github.com/TengdaHan/GAN-TensorFlow

 

[출처] http://solarisailab.com/archives/2482

 

 

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