Notebooks
M
Microsoft
AutoencodersTF

AutoencodersTF

artificial-intelligencernnganmicrosoft-for-beginnerslessonsAImicrosoft-AI-For-Beginnersmachine-learningdeep-learning09-Autoencoders4-ComputerVisioncomputer-visioncnnNLP

Autoencoders

When training CNNs, one of the problems is that we need a lot of labeled data. In the case of image classification, we need to separate images into different classes, which is a manual effort.

However, we might want to use raw (unlabeled) data for training CNN feature extractors, which is called self-supervised learning. Instead of labels, we will use training images as both network input and output. The main idea of autoencoder is that we will have an encoder network that converts input image into some latent space (normally it is just a vector of some smaller size), then the decoder network, whose goal would be to reconstruct the original image.

Since we are training autoencoder to capture as much of the information from the original image as possible for accurate reconstruction, the network tries to find the best embedding of input images to capture the meaning.

AutoEncoder Diagram

Image from Keras blog

Most of the examples below are inspired by this article.

Let's create simplest autoencoder for MNIST:

[1]
[2]
Output
[3]
[5]
[131]
Train on 60000 samples, validate on 10000 samples
Epoch 1/25
59648/60000 [============================>.] - ETA: 0s - loss: 0.2134
/usr/local/lib/python3.7/dist-packages/tensorflow/python/keras/engine/training.py:2325: UserWarning: `Model.state_updates` will be removed in a future version. This property should not be used in TensorFlow 2.0, as `updates` are applied automatically.
  warnings.warn('`Model.state_updates` will be removed in a future version. '
60000/60000 [==============================] - 6s 99us/sample - loss: 0.2130 - val_loss: 0.1454
Epoch 2/25
60000/60000 [==============================] - 5s 86us/sample - loss: 0.1353 - val_loss: 0.1258
Epoch 3/25
60000/60000 [==============================] - 5s 87us/sample - loss: 0.1225 - val_loss: 0.1177
Epoch 4/25
60000/60000 [==============================] - 5s 85us/sample - loss: 0.1163 - val_loss: 0.1126
Epoch 5/25
60000/60000 [==============================] - 5s 87us/sample - loss: 0.1120 - val_loss: 0.1091
Epoch 6/25
60000/60000 [==============================] - 5s 86us/sample - loss: 0.1093 - val_loss: 0.1070
Epoch 7/25
60000/60000 [==============================] - 5s 87us/sample - loss: 0.1072 - val_loss: 0.1055
Epoch 8/25
60000/60000 [==============================] - 5s 87us/sample - loss: 0.1057 - val_loss: 0.1041
Epoch 9/25
60000/60000 [==============================] - 5s 85us/sample - loss: 0.1045 - val_loss: 0.1028
Epoch 10/25
60000/60000 [==============================] - 5s 84us/sample - loss: 0.1035 - val_loss: 0.1022
Epoch 11/25
60000/60000 [==============================] - 5s 84us/sample - loss: 0.1026 - val_loss: 0.1011
Epoch 12/25
60000/60000 [==============================] - 5s 83us/sample - loss: 0.1018 - val_loss: 0.1003
Epoch 13/25
60000/60000 [==============================] - 5s 83us/sample - loss: 0.1012 - val_loss: 0.0996
Epoch 14/25
60000/60000 [==============================] - 5s 83us/sample - loss: 0.1005 - val_loss: 0.0991
Epoch 15/25
60000/60000 [==============================] - 5s 83us/sample - loss: 0.1000 - val_loss: 0.0988
Epoch 16/25
60000/60000 [==============================] - 5s 82us/sample - loss: 0.0995 - val_loss: 0.0981
Epoch 17/25
60000/60000 [==============================] - 5s 83us/sample - loss: 0.0990 - val_loss: 0.0976
Epoch 18/25
60000/60000 [==============================] - 5s 83us/sample - loss: 0.0986 - val_loss: 0.0974
Epoch 19/25
60000/60000 [==============================] - 5s 84us/sample - loss: 0.0982 - val_loss: 0.0969
Epoch 20/25
60000/60000 [==============================] - 5s 85us/sample - loss: 0.0978 - val_loss: 0.0970
Epoch 21/25
60000/60000 [==============================] - 5s 84us/sample - loss: 0.0975 - val_loss: 0.0962
Epoch 22/25
60000/60000 [==============================] - 5s 84us/sample - loss: 0.0971 - val_loss: 0.0960
Epoch 23/25
60000/60000 [==============================] - 5s 83us/sample - loss: 0.0968 - val_loss: 0.0958
Epoch 24/25
60000/60000 [==============================] - 5s 84us/sample - loss: 0.0966 - val_loss: 0.0953
Epoch 25/25
60000/60000 [==============================] - 5s 83us/sample - loss: 0.0963 - val_loss: 0.0953
<tensorflow.python.keras.callbacks.History at 0x7f3fa179b690>
[132]
/usr/local/lib/python3.7/dist-packages/tensorflow/python/keras/engine/training.py:2325: UserWarning: `Model.state_updates` will be removed in a future version. This property should not be used in TensorFlow 2.0, as `updates` are applied automatically.
  warnings.warn('`Model.state_updates` will be removed in a future version. '
OutputOutput
[133]
/usr/local/lib/python3.7/dist-packages/tensorflow/python/keras/engine/training.py:2325: UserWarning: `Model.state_updates` will be removed in a future version. This property should not be used in TensorFlow 2.0, as `updates` are applied automatically.
  warnings.warn('`Model.state_updates` will be removed in a future version. '
[134]
Output
[136]
6.3110805 0.0
Output

Task 1: Try to train autoencoder with very small latent vector size, eg. 2, and plot the dots corresponding to different digits. Hint: Use fully-connected dense layer after the convoluitonal part to reduce the vector size to the required value.

Task 2: Starting from different digits, obtain their latent space representations, and see what effect adding some noise to the latent space has on the resulting digits.

Denoising

Autoencoders can be effectively used to remove noise from images. In order to train denoiser, we will start with noise-free images, and add artificial noise to them. Then, we will feed autoencoder with noisy images as input, and noise-free images as output.

Let's see how this works for MNIST:

[137]
Output
[141]
Train on 60000 samples, validate on 10000 samples
Epoch 1/25
60000/60000 [==============================] - 6s 101us/sample - loss: 0.1576 - val_loss: 0.1566
Epoch 2/25
60000/60000 [==============================] - 6s 95us/sample - loss: 0.1564 - val_loss: 0.1553
Epoch 3/25
60000/60000 [==============================] - 6s 94us/sample - loss: 0.1555 - val_loss: 0.1539
Epoch 4/25
60000/60000 [==============================] - 6s 95us/sample - loss: 0.1545 - val_loss: 0.1530
Epoch 5/25
60000/60000 [==============================] - 6s 95us/sample - loss: 0.1538 - val_loss: 0.1517
Epoch 6/25
60000/60000 [==============================] - 6s 93us/sample - loss: 0.1528 - val_loss: 0.1506
Epoch 7/25
60000/60000 [==============================] - 6s 93us/sample - loss: 0.1521 - val_loss: 0.1499
Epoch 8/25
60000/60000 [==============================] - 5s 92us/sample - loss: 0.1514 - val_loss: 0.1495
Epoch 9/25
60000/60000 [==============================] - 6s 92us/sample - loss: 0.1508 - val_loss: 0.1487
Epoch 10/25
60000/60000 [==============================] - 6s 93us/sample - loss: 0.1500 - val_loss: 0.1483
Epoch 11/25
60000/60000 [==============================] - 6s 92us/sample - loss: 0.1495 - val_loss: 0.1484
Epoch 12/25
60000/60000 [==============================] - 6s 94us/sample - loss: 0.1487 - val_loss: 0.1468
Epoch 13/25
60000/60000 [==============================] - 6s 92us/sample - loss: 0.1482 - val_loss: 0.1467
Epoch 14/25
60000/60000 [==============================] - 6s 92us/sample - loss: 0.1476 - val_loss: 0.1459
Epoch 15/25
60000/60000 [==============================] - 6s 92us/sample - loss: 0.1469 - val_loss: 0.1450
Epoch 16/25
60000/60000 [==============================] - 6s 92us/sample - loss: 0.1463 - val_loss: 0.1442
Epoch 17/25
60000/60000 [==============================] - 6s 92us/sample - loss: 0.1457 - val_loss: 0.1441
Epoch 18/25
60000/60000 [==============================] - 6s 92us/sample - loss: 0.1451 - val_loss: 0.1429
Epoch 19/25
60000/60000 [==============================] - 6s 92us/sample - loss: 0.1445 - val_loss: 0.1425
Epoch 20/25
60000/60000 [==============================] - 6s 94us/sample - loss: 0.1440 - val_loss: 0.1418
Epoch 21/25
60000/60000 [==============================] - 6s 93us/sample - loss: 0.1435 - val_loss: 0.1423
Epoch 22/25
60000/60000 [==============================] - 6s 93us/sample - loss: 0.1430 - val_loss: 0.1409
Epoch 23/25
60000/60000 [==============================] - 6s 94us/sample - loss: 0.1426 - val_loss: 0.1405
Epoch 24/25
60000/60000 [==============================] - 6s 93us/sample - loss: 0.1422 - val_loss: 0.1409
Epoch 25/25
60000/60000 [==============================] - 6s 93us/sample - loss: 0.1418 - val_loss: 0.1398
<tensorflow.python.keras.callbacks.History at 0x7f3fa612c4d0>
[144]
OutputOutput

Exercise: See how denoiser trained on MNIST digits works for different images. As an example, you can take Fashion MNIST dataset, which has the same image size. Note that denoiser works well only on the same image type that it was trained on (i.e. for the same probability distribution of input data).

Super-resolution

Similarly to denoiser, we can train autoencoders to increase the resolution of the image. To train super-resolution network, we will start with high-resolution images, and automatically downscale them to produce network inputs. We will then feed autoencoder with small images as inputs and high-res images as outputs.

Let's downscale MNIST to 14x14:

[6]
Output
[7]
[8]
Epoch 1/25
469/469 [==============================] - 6s 10ms/step - loss: 0.3413 - val_loss: 0.1519
Epoch 2/25
469/469 [==============================] - 4s 9ms/step - loss: 0.1457 - val_loss: 0.1292
Epoch 3/25
469/469 [==============================] - 4s 9ms/step - loss: 0.1273 - val_loss: 0.1202
Epoch 4/25
469/469 [==============================] - 4s 9ms/step - loss: 0.1189 - val_loss: 0.1142
Epoch 5/25
469/469 [==============================] - 4s 9ms/step - loss: 0.1148 - val_loss: 0.1107
Epoch 6/25
469/469 [==============================] - 4s 9ms/step - loss: 0.1115 - val_loss: 0.1083
Epoch 7/25
469/469 [==============================] - 4s 9ms/step - loss: 0.1093 - val_loss: 0.1063
Epoch 8/25
469/469 [==============================] - 4s 9ms/step - loss: 0.1071 - val_loss: 0.1046
Epoch 9/25
469/469 [==============================] - 4s 9ms/step - loss: 0.1060 - val_loss: 0.1037
Epoch 10/25
469/469 [==============================] - 4s 9ms/step - loss: 0.1048 - val_loss: 0.1026
Epoch 11/25
469/469 [==============================] - 4s 9ms/step - loss: 0.1039 - val_loss: 0.1019
Epoch 12/25
469/469 [==============================] - 4s 9ms/step - loss: 0.1030 - val_loss: 0.1012
Epoch 13/25
469/469 [==============================] - 4s 9ms/step - loss: 0.1024 - val_loss: 0.1004
Epoch 14/25
469/469 [==============================] - 4s 9ms/step - loss: 0.1017 - val_loss: 0.0999
Epoch 15/25
469/469 [==============================] - 4s 9ms/step - loss: 0.1010 - val_loss: 0.0993
Epoch 16/25
469/469 [==============================] - 4s 9ms/step - loss: 0.1005 - val_loss: 0.0989
Epoch 17/25
469/469 [==============================] - 4s 9ms/step - loss: 0.0999 - val_loss: 0.0983
Epoch 18/25
469/469 [==============================] - 4s 9ms/step - loss: 0.0995 - val_loss: 0.0982
Epoch 19/25
469/469 [==============================] - 4s 9ms/step - loss: 0.0990 - val_loss: 0.0975
Epoch 20/25
469/469 [==============================] - 4s 9ms/step - loss: 0.0987 - val_loss: 0.0971
Epoch 21/25
469/469 [==============================] - 4s 9ms/step - loss: 0.0981 - val_loss: 0.0971
Epoch 22/25
469/469 [==============================] - 4s 9ms/step - loss: 0.0979 - val_loss: 0.0965
Epoch 23/25
469/469 [==============================] - 4s 9ms/step - loss: 0.0977 - val_loss: 0.0959
Epoch 24/25
469/469 [==============================] - 4s 9ms/step - loss: 0.0972 - val_loss: 0.0957
Epoch 25/25
469/469 [==============================] - 4s 9ms/step - loss: 0.0972 - val_loss: 0.0955
<tensorflow.python.keras.callbacks.History at 0x7f66790ada90>
[9]
OutputOutput

Exercise: Try to train super-resolution network on CIFAR-10 for 2x and 4x upscaling. Use noise as input to 4x upscaling model and observe the result.

Variational Auto-Encoders (VAE)

Traditional autoencoders reduce the dimension of the input data somehow, figuring out the important features of input images. However, latent vectors often do not make much sense. In other words, taking MNIST dataset as an example, figuring out which digits correspond to different latent vectors is not an easy task, because close latent vectors would not necessarily correspond to the same digits.

On the other hand, to train generative models it is better to have some understanding of the latent space. This idea leads us to variational auto-encoder (VAE).

VAE is the autoencoder that learns to predict statistical distribution of the latent parameters, so-called latent distribution. For example, we can assume that latent vectors would be distributed as N(z_mean,ez_log_sigma)N(\mathrm{z\_mean},e^{\mathrm{z\_log\_sigma}}), where z_mean,z_log_sigma∈Rd\mathrm{z\_mean}, \mathrm{z\_log\_sigma} \in\mathbb{R}^d. Encoder in VAE learns to predict those parameters, and then decoder takes a random vector from this distribution to reconstruct the object.

To summarize:

  • From input vector, we predict z_mean and z_log_sigma (instead of predicting the standard deviation itself, we predict it's logarithm)
  • We sample a vector sample from the distribution N(z_mean,ez_log_sigma)N(\mathrm{z\_mean},e^{\mathrm{z\_log\_sigma}})
  • Decoder tries to decode the original image using sample as an input vector
[21]
[22]
[23]

Variational auto-encoders use complex loss function that consists of two parts:

  • Reconstruction loss is the loss function that shows how close reconstructed image is to the target (can be MSE). It is the same loss function as in normal autoencoders.
  • KL loss, which ensures that latent variable distributions stays close to normal distribution. It is based on the notion of Kullback-Leibler divergence - a metric to estimate how similar two statistical distributions are.
[24]
[25]
Train on 60000 samples, validate on 10000 samples
Epoch 1/25
59520/60000 [============================>.] - ETA: 0s - loss: 48.6396
/usr/local/lib/python3.7/dist-packages/tensorflow/python/keras/engine/training.py:2325: UserWarning: `Model.state_updates` will be removed in a future version. This property should not be used in TensorFlow 2.0, as `updates` are applied automatically.
  warnings.warn('`Model.state_updates` will be removed in a future version. '
60000/60000 [==============================] - 4s 64us/sample - loss: 48.5874 - val_loss: 41.8877
Epoch 2/25
60000/60000 [==============================] - 3s 57us/sample - loss: 41.1296 - val_loss: 40.2556
Epoch 3/25
60000/60000 [==============================] - 3s 56us/sample - loss: 40.0063 - val_loss: 39.3692
Epoch 4/25
60000/60000 [==============================] - 3s 56us/sample - loss: 39.2531 - val_loss: 38.7666
Epoch 5/25
60000/60000 [==============================] - 3s 57us/sample - loss: 38.7147 - val_loss: 38.6124
Epoch 6/25
60000/60000 [==============================] - 3s 57us/sample - loss: 38.2962 - val_loss: 38.1867
Epoch 7/25
60000/60000 [==============================] - 3s 56us/sample - loss: 37.9756 - val_loss: 37.9831
Epoch 8/25
60000/60000 [==============================] - 3s 57us/sample - loss: 37.6933 - val_loss: 37.5475
Epoch 9/25
60000/60000 [==============================] - 3s 57us/sample - loss: 37.4323 - val_loss: 37.2913
Epoch 10/25
60000/60000 [==============================] - 3s 56us/sample - loss: 37.2133 - val_loss: 37.1992
Epoch 11/25
60000/60000 [==============================] - 3s 57us/sample - loss: 36.9966 - val_loss: 36.9521
Epoch 12/25
60000/60000 [==============================] - 3s 57us/sample - loss: 36.8204 - val_loss: 36.8431
Epoch 13/25
60000/60000 [==============================] - 3s 57us/sample - loss: 36.6490 - val_loss: 36.6979
Epoch 14/25
60000/60000 [==============================] - 3s 57us/sample - loss: 36.5023 - val_loss: 36.6661
Epoch 15/25
60000/60000 [==============================] - 3s 57us/sample - loss: 36.3456 - val_loss: 36.4957
Epoch 16/25
60000/60000 [==============================] - 3s 56us/sample - loss: 36.2266 - val_loss: 36.6669
Epoch 17/25
60000/60000 [==============================] - 3s 57us/sample - loss: 36.1045 - val_loss: 36.4855
Epoch 18/25
60000/60000 [==============================] - 3s 57us/sample - loss: 35.9922 - val_loss: 36.4150
Epoch 19/25
60000/60000 [==============================] - 3s 56us/sample - loss: 35.8968 - val_loss: 36.1196
Epoch 20/25
60000/60000 [==============================] - 3s 57us/sample - loss: 35.7991 - val_loss: 36.0708
Epoch 21/25
60000/60000 [==============================] - 3s 56us/sample - loss: 35.7129 - val_loss: 36.1686
Epoch 22/25
60000/60000 [==============================] - 3s 57us/sample - loss: 35.6214 - val_loss: 36.1080
Epoch 23/25
60000/60000 [==============================] - 3s 57us/sample - loss: 35.5357 - val_loss: 36.2309
Epoch 24/25
60000/60000 [==============================] - 3s 56us/sample - loss: 35.4528 - val_loss: 36.1416
Epoch 25/25
60000/60000 [==============================] - 3s 56us/sample - loss: 35.3650 - val_loss: 35.7258
<tensorflow.python.keras.callbacks.History at 0x7f0e00233890>
[26]
/usr/local/lib/python3.7/dist-packages/tensorflow/python/keras/engine/training.py:2325: UserWarning: `Model.state_updates` will be removed in a future version. This property should not be used in TensorFlow 2.0, as `updates` are applied automatically.
  warnings.warn('`Model.state_updates` will be removed in a future version. '
OutputOutput
[27]
/usr/local/lib/python3.7/dist-packages/tensorflow/python/keras/engine/training.py:2325: UserWarning: `Model.state_updates` will be removed in a future version. This property should not be used in TensorFlow 2.0, as `updates` are applied automatically.
  warnings.warn('`Model.state_updates` will be removed in a future version. '
Output
[28]
/usr/local/lib/python3.7/dist-packages/tensorflow/python/keras/engine/training.py:2325: UserWarning: `Model.state_updates` will be removed in a future version. This property should not be used in TensorFlow 2.0, as `updates` are applied automatically.
  warnings.warn('`Model.state_updates` will be removed in a future version. '
Output

Task: In our sample, we have trained fully-connected VAE. Now take the CNN from traditional auto-encoder above and create CNN-based VAE.