Three lessons from our AI projects
Over two years, we supported dozens of AI projects, from initial exploration to a model that runs every night. Some went exactly as planned; others…
Generating people, animals, b&b’s, or art with AI is one of many interesting developments in machine learning. The neural networks used on these websites belong to a category called Generative Adversarial Networks, or GANs. The idea behind GANs was developed in 2014 by Ian Goodfellow, then a research scientist at Google. In his paper, he compares how a GAN works to the cat-and-mouse game between counterfeiters and the police. The counterfeiters aim to reproduce money as convincingly as possible, while the police try to detect as much counterfeit money as they can. By identifying counterfeit money, the police get better at detecting it. The ‘feedback’ the counterfeiters receive from this, in turn, helps them produce better counterfeits. This continues until the counterfeit money can no longer be distinguished from the real thing. A GAN’s ability to generate comes from the counterfeiters, but it depends on feedback from the police. In a GAN, the counterfeiters and the police are two neural networks that learn from each other in a feedback loop.
Beyond generating images, GANs can also be used for other practical purposes. Examples include: [removing visual noise (such as rain) from images, upscaling low-resolution images, restoring damaged photos, generating a personalised emoji from a photo of someone (e.g. a Bitmoji on Snapchat), expanding datasets, and many other applications. Although this blog focuses on the generative power of GANs, it is worth noting that GANs can do much more than generate images. If you enjoy making visual things but do not consider yourself creative, these networks are ideal. The most important ingredient is thousands of images to train the network on. For my first dive into the world of GAN’s, the data came from the Kaggle competition Generative Dog Images. Alongside the data, the discussions there were also a useful source of inspiration while developing a GAN. But before diving into the code, it helps to understand what these almost magical neural networks look like under the bonnet.
What is a GAN? Put simply, a GAN is a neural network that learns the distribution of a given dataset, making it possible to generate new samples from it. This blog focuses on generating images, but GANs can generate anything from text to audio. A GAN’s architecture has two parts: the generator and the discriminator. The generator is a network that produces a fake sample from the dataset based on a ‘noise vector’. The noise vector (z) is a one-dimensional vector of arbitrary length containing random numbers. It is often 100 elements long, with the numbers drawn from a standard normal distribution. The discriminator is a network that works like a standard binary classifier. During training, it is shown real samples from the dataset and fake, generated samples, which it then classifies as real or fake. Based on that feedback, the generator adjusts its weights to produce better samples that the discriminator will hopefully classify as ‘real’. A GAN trains itself by alternating between generation and classification. In pseudocode, the process looks like this:
Here you can see that the GAN is trained in two steps. The first step trains the discriminator, and only the discriminator’s weights are updated. The generator stays the same during this step. In the second step, the generator produces new samples, which the discriminator classifies. That classification is used to calculate the GAN’s loss and update the generator’s weights. The backpropagation that passes through the discriminator to the generator in this step is also why the GAN is one network, despite consisting of two networks. They must be connected for the generator’s weights to be updated. To get a little more technical, we can define the discriminator (D) and the generator (G) as differentiable functions. The function D(x) represents the probability that input x is a real sample. The network’s aim is to maximise the probability that this function classifies samples correctly. Meanwhile, the function G(z) maps a noise vector z to a data point, with the aim of minimising 1 - D(G(z)) — or, in plain language, minimising the probability that D sees a generated sample as fake. This interplay between D and G is a minimax game and can be expressed by the following function:
Source
So far, we have not described the architectures of D and G in detail. That is because these functions look different for each type of GAN. What forms can these functions take, and how does that affect the GAN’s output?
GAN architectures
As with any neural network, many architectures are possible, and GANs are no different. They range from the original GAN, which used only multilayer perceptrons, to BigGAN , which uses techniques such as attention maps, skip-z connections and more. As the image shows, the quality of GAN output has improved considerably. Despite the large difference in quality, both networks still have the same generator/discriminator structure; BigGAN simply has a much more advanced discriminator and generator.
Output from the original GAN (left) and BigGAN (right)
Another example of a network that produces high-quality output is StyleGAN. This network also has the classic generator–discriminator structure, but a specialised training method enables it to produce high-resolution images. You can find more information about other generative networks, including StyleGAN, in this blog. Somewhere between the original GAN and BigGAN is the Deep Convolutional GAN (DCGAN). This network uses convolutional layers in both the discriminator and the generator. When trained well, its relatively simple architecture can still produce good output. The DCGAN’s generator and discriminator are both convolutional networks, with the discriminator mirroring the generator. The image below shows the architecture of the generator network. The discriminator mirrors the generator, except that it ends with a sigmoid node that predicts whether the input is real or fake. An important detail is that DCGAN does not use pooling or upscaling layers, although these are common in convolutional networks. Instead, it uses (de)convolutions with a stride (step size) of 2. This allows the network to learn the transformations, producing better output. The creators of DCGAN also recommend:
Architecture of the DCGAN generator
With these guidelines and an idea of how to structure the networks, it might sound as though training them is straightforward. Unfortunately, that is not the case: GANs are notoriously difficult to train. What can go wrong, and how do you fix it?
How do you train a GAN?
Training a GAN is a delicate process, and a lot can go wrong. The aim is to keep the discriminator and generator in balance: neither should become “stronger” than the other. Ideally, the generator and discriminator reach a Nash equilibrium. This term from game theory describes a situation in which a player does not change their strategy, regardless of what the other player does. Reaching this equilibrium is very difficult. More often, you will see the advantage shift back and forth between the players. That is not a problem: a GAN can still generate good output in this situation.
Checkerboard artefacts (left), [colourful grids (centre), mode collapse (right)
Some things that can go wrong when training a GAN are shown above. These were also the three biggest problems I encountered when developing a GAN. Checkerboard artefacts are a checkerboard-like pattern of pixels, or groups of pixels, that are darker or lighter than their neighbours. These artefacts arise during deconvolution. If the kernel size and stride are combined incorrectly, the learned filter is applied to the same pixel more than once, changing its value twice as much as those of its neighbours. Fortunately, there is a simple solution: use a kernel size that is a multiple of the stride. This makes the filter pass over each pixel exactly once, preventing the artefacts. If this combination of stride and kernel size is not possible, another option is to use a stride of 1 for the filter in the final convolutional layer. This also reduces the artefacts. You can find more information about these artefacts in this article by Distill. If a GAN’s output looks like a colourful grid, the learning rates of the generator and discriminator are probably too far out of balance. A GAN produces output like this when the generator learns too quickly or the discriminator learns too slowly. You can also see this in the generator’s loss: if it fluctuates sharply, the learning rates need adjusting. The discriminator’s and generator’s losses should follow roughly the same pattern, with the generator’s loss usually slightly higher than the discriminator’s. The biggest problem when training a GAN is (partial) mode collapse. In this case, the GAN’s output has little variation and everything looks alike. You can see this in the image on the right, for example, where the model generates a 6 for every possible input value of z. This happens when the generator ‘discovers’ that the discriminator always classifies one particular output as real. The generator then receives positive feedback and keeps generating that output. This makes sense, because the generator’s goal is to fool the discriminator. Researchers have not yet established the precise cause of mode collapse, but they have proposed many solutions. The simplest is to add noise to the discriminator’s labels; more on that later. Another, more complex way to counter mode collapse is the unrolled GAN. In this architecture, each training step updates the generator using not only the current discriminator but also the discriminator as it would be N steps into the future. This better discriminator output helps the generator learn more effectively and keeps its output diverse. A few other small adjustments to the GAN that helped considerably were flipping the labels fed to the discriminator and adding noise to those same labels. Swapping the labels (1 for fake, 0 for real) helps the generator learn faster in the early stages of training. Adding noise means, first, that the labels are ‘soft’: rather than fixed values of 1 or 0, they take random values between 1 and 0.9 or between 0 and 0.1. Second, the labels had a 5% chance of being flipped, so a 1 became a 0 and vice versa. This reduces the likelihood of mode collapse and was enough to prevent it for the datasets used. Finally, batch normalization was used in the discriminator, and spectral normalization in both the discriminator and the generator. Batch normalisation makes training in deep networks more stable, while spectral normalisation produces better, sharper output.
With all this knowledge and these tips in hand, it was time to get down to the real work. First, the GAN was tested on the ‘hello world’ of machine learning datasets: MNIST digits. This is a relatively simple dataset to start with because it is black and white and has little variation. To stay consistent with the DCGAN paper, the data was scaled up from 28 by 28 to 64 by 64 pixels. MNIST is a good dataset to start with: it took little effort to produce clear output, as you can see below.
Because the MNIST dataset is such an obvious choice and not particularly interesting, the GAN was also trained on drawings of sheep from Google’s Quick, draw! dataset. The idea came from the book ‘Dreaming of Electric Sheep’, which contains 10,000 GAN-generated drawings of sheep. Below is an image showing drawings from the dataset alongside those generated by the GAN.
Drawings from the Quick, Draw dataset (left) and generated drawings (right)
As you can see, the GAN generates black-and-white data well. Things get more interesting when we train it on colour images: 3D data with multiple colour channels. The extra colour dimension adds complexity, making it much harder for the GAN to generate realistic data. Choosing hyperparameters is also more difficult, as many parameter combinations produce the colourful grids discussed earlier. A simple colour dataset to start with is the Google StreetView House Number (SVHN) dataset. It contains close-up images of house numbers that vary in colour, font and camera angle. Getting good results on this dataset required considerable adjustments to the learning rates of the discriminator and generator. In the end, learning rates of 0.001 and 0.0001 for the discriminator and generator, respectively, worked best on the colour datasets. A sample from the dataset and the generated images are shown below.
SVHN dataset (left) and generated samples (right)
[divider line_type="No Line" custom_height="30″]
After gradually warming up on these practice datasets, it was time to move on to the dataset the GAN was ultimately built for: the Stanford Dogs dataset, which was the focus of the Kaggle competition. This dataset is much more complex than the SVHN data. It contains a wide variety of dogs, differing in colour, breed and posture, in a wide variety of settings. Unfortunately, this complexity proved too difficult for the network to learn, and its output is easy to distinguish from real images. Even so, the output does bear some resemblance to a dog. After days of training and experimenting with learning rates, we eventually produced a network that gave good output. The final network took 9 hours to train, with learning rates of 0.0001 and 0.0003 for the discriminator and generator, respectively. A sample of the network’s output is shown below.
Some dogs generated by the GAN
Interestingly, the GAN can represent some poses well, such as a head looking upwards or a side view of a dog. It was also interesting to see that it had learnt the colours and patterns of the coat well, with good variation in both.
The GAN’s output is good enough for you to recognise dogs, but there is still plenty of room for improvement. As mentioned earlier, the Stanford Dogs dataset is complex because the data contains a lot of variation. One possible solution is to use a Conditional GAN (CGAN). In a CGAN, the generator and discriminator receive a class label alongside their usual input. These class labels are provided in a form such as an embedding or a one-hot vector. A fully connected layer then converts the class labels to the right format before they are combined with the original input. The combined input can then be passed through the discriminator or generator. For the Stanford Dogs dataset, the class labels could, for example, represent the dog’s breed. They could also represent a completely different feature of the image, such as the dog’s colour, the angle from which the photo was taken, the background, etc.
One advantage of a CGAN is that the labels help the network distinguish between variations in the dataset, resulting in higher-quality output. The CGAN architecture also makes it possible to generate output from just one class, which is very difficult with a standard GAN. The main disadvantage of using a CGAN is that you move from an unsupervised to a supervised problem, so you need a labelled dataset.
Another solution is to use a Wasserstein GAN (WGAN). A WGAN replaces the discriminator with a so-called critic. Instead of making a hard ‘real or fake’ prediction, the critic assigns a score for how real the input is. This score should be low for real input and high for fake input, and is not bounded between 0 and 1. The WGAN also uses a new loss function: Wasserstein loss. Exactly how Wasserstein loss works in a WGAN is too complex to explain briefly here; you can find a detailed explanation here. One advantage of a WGAN is that you can optimise the critic without having to account for what that means for the generator. In a standard GAN, a perfect discriminator means the generator learns nothing, but with a perfect critic, the generator can still improve. You can see this in the image below. It compares the gradients of a perfect critic and discriminator when distinguishing between two (hypothetical) distributions. As you can see, the discriminator’s sigmoid function produces poor gradients, while the critic’s linear activation produces a gradient that the generator can still learn from.
This property of the critic means that a WGAN no longer needs to maintain a balance between the generator and critic, making it more robust and less likely to fail. This robustness also allows more variation in the architecture and learning rates of the generator and critic/discriminator without the output becoming unrecognisable. A final advantage is that a WGAN can also prevent mode collapse.
Of course, the best way to get high-quality output is to use a state-of-the-art architecture, such as the previously mentioned StyleGAN or BigGAN. These architectures are very complex to implement yourself (though not impossible). Training these networks also requires good hardware: StyleGAN, for example, needed one week on 8 GPUs. For a ‘real’ project, however, these solutions are worth considering.
After reading many papers and blog posts about developing GANs, we managed to build a working version of DCGAN. This network was able to capture the distribution of relatively low-complexity datasets and generate new data from it. Unfortunately, the Stanford Dogs dataset was too complex for the current implementation of the network.
Despite all the research that went into development, it soon became clear that there was no one-size-fits-all solution for every dataset; the learning rates always needed some adjustment before the output was satisfactory. In some cases, a change of 0.0001 made the difference between good output and colourful grids. This made training a GAN difficult, but it was also a useful exercise in reading and interpreting the network’s losses.
Researching and developing a GAN was a valuable learning experience, and it was exciting to see that a relatively simple network could still generate good results. Unfortunately, the GAN had not managed to learn the anatomy of a dog well, as the output shows. Even so, the results were good enough to place in the top 47% of the Kaggle competition, and when the image above was pasted into Word, it was given the description “Image with dog”. The dogs may not be good enough to fool a human, but they do seem to fool Microsoft’s algorithms.
Over two years, we supported dozens of AI projects, from initial exploration to a model that runs every night. Some went exactly as planned; others…
Many organisations have now run an AI pilot. The model works, the demo gets applause, and then nothing else happens. In our experience, most projects…
Artificial Intelligence is developing rapidly. New models appear almost every week, and more and more organisations are experimenting with AI. At the…
Want to be the first to hear about a new blog post?
Thanks for signing up!