1 of 41

WGAN-GP

Alex Lau

Hong Kong Machine Learning Meetup

17/07/2019

2 of 41

Table of Content

  • What is GAN
  • Pain-points when training GAN
  • What is WGAN-GP
  • How WGAN-GP can help
  • WGAN-GP v.s. GAN
  • WGAN-GP Applications

3 of 41

What is GAN?

  • Proposed by Ian Goodfellow in 2014
  • Generative Model
    • Input training data set (aka real data)
    • Output data highly resemble real data
    • Statistically, it is trying to approximate real data distribution

4 of 41

What is GAN?

Discriminator

Generator

Input

Data (real / fake)

Data (prior noise)

Output

Probability (0-1)

~1 means high confidence the input data belong to real data

Fake data

Purpose (Functional)

Classify fake data and real data correctly

Fool discriminator to mistake fake data as real data

Purpose (Math)

Maximise the chance of distinguishing real data as real data

Minimise the chance of classifying fake data as fake data

5 of 41

How do we measure distribution similarity?

  • Important metrics!
  • Real data distribution v.s. Fake data distribution
  • GAN model measure distribution similarity by JS (Jensen-Shannon) divergence
    • computed by KL (Kullback Leibler) divergence
    • Both measures are commonly used in machine learning

6 of 41

Cost Function of GAN

D stands for Discriminator network

G stands for Generator network

First term considers sample from real data

  1. Sample from real data, (x1, x2 ….xn)
  2. Input sample into discriminator
  3. Collect discriminator outputs, D(x1), D(x2)... D(xn)
  4. Take log on the outputs, log(D(x1)) … log(D(xn))
  5. Take average of the collections

Second term consider sample from fake data

  1. Sample prior noise
  2. Feed to generator to get fake data
  3. Feed fake data into discriminator
  4. ...

7 of 41

Cost Function of GAN

D stands for Discriminator network

G stands for Generator network

D(x) should output ~1 when input are real data

For sample from real data, discriminator aims to higher the chance of classifying them as real data

D(x) should output ~0 when input are generated data

For fake data, discriminator aims to higher the chance of classifying them as fake data

8 of 41

Cost Function of GAN

D stands for Discriminator network

G stands for Generator network

G(.) should fool D(.) so as to output ~1

Generator aims to produce data so as to higher the chance of discriminator classifying them as real data

9 of 41

Cost Function of GAN

D stands for Discriminator network

G stands for Generator network

Optimizing the cost function is equivalent to minimizing JS divergence

10 of 41

Pain-point when training GAN

11 of 41

Cost Function not Smooth

GAN uses JS divergence (closely related to KL divergence) as measurement of distribution similarity

Both JS divergence and KL divergence gives non-smooth surface when two distribution has a disjoint support

12 of 41

Vanishing Gradient

If we apply KL/ JS divergence on the cost function, it easily leads to vanishing gradient

(Vanishing gradient is an issue when optimise comes to a flat surface...

It makes neural network learn almost nothing for each iteration)

13 of 41

Mode Collapse

Generated data stuck in a small space with extremely low variety

Equilibrium is reached but the model is bad...

14 of 41

Lack of Indicative Metrics for Training Performance

15 of 41

Lack of Indicative Metrics for Training Performance

16 of 41

What is WGAN-GP?

  • Variant of GAN proposed in 2017, come after WGAN (both 1000+ citation)
  • Wasserstein GAN with Gradient Penalty
  • What’s new?
    • New cost function with Wasserstein distance
    • Enforce optimization constraint with gradient penalty

17 of 41

Introducing Wasserstein Distance

  • Another measurement on distribution similarity
  • It even works well when two distributions have disjoint support (overcome the difficulty faced by JS divergence)

18 of 41

How Wasserstein Distance Measure Similarity?

  • View as an optimal transport problem
  • Blue: original configuration
  • New: target configuration
  • How much effort at minimum I need to take to make original configuration into a target configuration
  • Factors
    • Transport distance
    • Mass to be transport

v.s.

19 of 41

How Wasserstein Distance Measure Similarity?

Wasserstein Distance

= * (5-1) + * (6-2) + * (7-3) = 4

5

6

7

1

1/3

0

0

2

0

1/3

0

3

0

0

1/3

20 of 41

However, WD is Hard to Compute ….

Original form is VERY hard to compute

With math magic, it can be expressed in another form that is easier to compute (under a certain constraint…)

Apply gradient penalty to enforce the constraint (help avoid gradient exploding/ vanishing)

(Where D(x) is 1-Lipschitz function)

I.e. differentiable + derivative norm <= 1

21 of 41

How WGAN-GP Can Help?

22 of 41

Smoothen the Function to be Optimised

  • Way smoother
  • Less flat surface

23 of 41

Avoid Mode Collapse

  • Cost function with wasserstein distance is
    • More smooth
    • Less flat surface
  • Optimisation works better
  • Generator can learn effectively

24 of 41

Loss Curve More indicative of Training Performance

  • Loss curve can be a reliable indicator to tell if the training performance is going well
  • It goes well when generator’s loss and discriminator’s loss convergence and close to each other

25 of 41

Loss Curve More indicative of Training Performance

26 of 41

WGANGP v.s. GAN

27 of 41

GAN v.s. WGAN-GP in Algorithm

WGAN-GP

GAN

28 of 41

GAN v.s. WGAN-GP in Algorithm

WGAN-GP

GAN

What are somethings in common?

High-level structure

29 of 41

GAN v.s. WGAN-GP in Algorithm

WGAN-GP

GAN

What are somethings in common?

Generator input

30 of 41

GAN v.s. WGAN-GP in Algorithm

WGAN-GP

GAN

What are somethings in common?

Number of training for discriminator

31 of 41

GAN v.s. WGAN-GP in Algorithm

WGAN-GP

GAN

What are the differences?

Number of training for discriminator

32 of 41

GAN v.s. WGAN-GP in Algorithm

WGAN-GP

GAN

What are the differences?

Optimizer

  1. Minibatch
  2. Take into account momentum (i.e. preceding gradient direction)
  3. Take into account diminishing learning rate
  • Minibatch

33 of 41

GAN v.s. WGAN-GP in Algorithm

WGAN-GP

GAN

What are the differences?

Discriminator output

Discriminator output probability (0 - 1)

Sigmoid as activation function

Critic output any real number

Sigmoid is removed

34 of 41

GAN v.s. WGAN-GP in Algorithm

WGAN-GP

GAN

What are the differences?

Cost Function

35 of 41

Applications of WGAN-GP

36 of 41

Image Inpainting (Yu et al, 2018)

37 of 41

Image Inpainting (Yu et al, 2018)

Recovered by WGAN-GP

Recovered by GAN

38 of 41

Image Inpainting (Yu et al, 2018)

You can try their demo: http://jiahuiyu.com/deepfill/

39 of 41

Image Coloring

40 of 41

Reference

41 of 41

Open to Project Collaboration

Proposed Projects

  • Counting densely crowd
  • DeepFake