Variational Memory Addressing in Generative Models

Jörg Bornschein, Andriy Mnih, Daniel Zoran, Danilo J. Rezende

Introduction

Recent years have seen rapid developments in generative modelling. Much of the progress was driven by the use of powerful neural networks to parameterize conditional distributions composed to define the generative process (e.g., VAEs (kingma2013vae, ; rezende2014stochastic, ), GANs (goodfellow2014generative, )). In the Variational Autoencoder (VAE) framework for example, we typically define a generative model p(z)p({\mathbf{z}}), pθ(x∣z)p_{\theta}({\mathbf{x}}|{\mathbf{z}}) and an approximate inference model qϕ(z∣x)q_{\phi}({\mathbf{z}}|{\mathbf{x}}). All conditional distributions are parameterized by multilayered perceptrons (MLPs) which, in the simplest case, output the mean and the diagonal variance of a Normal distribution given the conditioning variables. We then optimize a variational lower bound to learn the generative model for x{\mathbf{x}}. Considering recent progress, we now have the theory and the tools to train powerful, potentially non-factorial parametric conditional distributions p(x∣y)p({\mathbf{x}}|{\mathbf{y}}) that generalize well with respect to x{\mathbf{x}} (normalizing flows (rezende2015flows, ), inverse autoregressive flows (kingma2016iaf, ), etc.).

Another line of work which has been gaining popularity recently is memory augmented neural networks (Das92learningcontextfree, ; Sukhbaatar2015endmemnet, ; graves2016hybrid, ). In this family of models the network is augmented with a memory buffer which allows read and write operations and is persistent in time. Such models usually handle input and output to the memory buffer using differentiable “soft” write/read operations to allow back-propagating gradients during training.

Here we propose a memory-augmented generative model that uses a discrete latent variable aa acting as an address into the memory buffer M{\bf M}. This stochastic perspective allows us to introduce a variational approximation over the addressing variable which takes advantage of target information when retrieving contents from memory during training. We compute the sampling distribution over the addresses based on a learned similarity measure between the memory contents at each address and the target. The memory contents ma{\mathbf{m}}_{a} at the selected address aa serve as a context for a continuous latent variable z{\mathbf{z}}, which together with ma{\mathbf{m}}_{a} is used to generate the target observation. We therefore interpret memory as a non-parametric conditional mixture distribution. It is non-parametric in the sense that we can change the content and the size of the memory from one evaluation of the model to another without having to relearn the model parameters. And since the retrieved content ma{\mathbf{m}}_{a} is dependent on the stochastic variable aa, which is part of the generative model, we can directly use it downstream to generate the observation x{\mathbf{x}}. These two properties set our model apart from other work on VAEs with mixture priors (dilokthanakul2016deep, ; nalisnick2016approximate, ) aimed at unconditional density modelling. Another distinguishing feature of our approach is that we perform sampling-based variational inference on the mixing variable instead of integrating it out as is done in prior work, which is essential for scaling to a large number of memory addresses.

Most existing memory-augmented generative models use soft attention with the weights dependent on the continuous latent variable to access the memory. This does not provide clean separation between inferring the address to access in memory and the latent factors of variation that account for the variability of the observation relative to the memory contents (see Figure 1). Or, alternatively, when the attention weights depend deterministically on the encoder, the retrieved memory content can not be directly used in the decoder.

Our contributions in this paper are threefold: a) We interpret memory-read operations as conditional mixture distribution and use amortized variational inference for training; b) demonstrate that we can combine discrete memory addressing variables with continuous latent variables to build powerful models for generative few-shot learning that scale gracefully with the number of items in memory; and c) demonstrate that the KL divergence over the discrete variable aa serves as a useful measure to monitor memory usage during inference and training.

Model and Training

We will now describe the proposed model along with the variational inference procedure we use to train it. The generative model has the form

where x{\mathbf{x}} is the observation we wish to model, aa is the addressing categorical latent variable, z{\mathbf{z}} the continuous latent vector, M{\bf M} the memory buffer and ma{\mathbf{m}}_{a} the memory contents at the aath address.

The generative process proceeds by first sampling an address aa from the categorical distribution p(a∣M)p(a|{\bf M}), retrieving the contents ma{\mathbf{m}}_{a} from the memory buffer M{\bf M}, and then sampling the observation xx from a conditional variational auto-encoder with ma{\mathbf{m}}_{a} as the context conditioned on (Figure 1, B). The intuition here is that if the memory buffer contains a set of templates, a trained model of this type should be able to produce observations by distorting a template retrieved from a randomly sampled memory location using the conditional variational autoencoder to account for the remaining variability.

We can write the variational lower bound for the model in (1):

In the rest of the paper, we omit the dependence on M{\bf M} for brevity. We will now describe the components of the model and the variational posterior (3) in detail.

The first component of the model is the memory buffer M{\bf M}. We here do not implement an explicit write operation but consider two possible sources for the memory content: Learned memory: In generative experiments aimed at better understanding the model’s behaviour we treat M{\bf M} as model parameters. That is we initialize M{\bf M} randomly and update its values using the gradient of the objective. Few-shot learning: In the generative few-shot learning experiments, before processing each minibatch, we sample ∣M∣|{\bf M}| entries from the training data and store them in their raw (pixel) form in M{\bf M}. We ensure that the training minibatch {x1,...,x∣B∣}\{{\mathbf{x}}_{1},...,{\mathbf{x}}_{|{\cal B}|}\} contains disjoint samples from the same character classes, so that the model can use M{\bf M} to find suitable templates for each target x{\mathbf{x}}.

The second component is the addressing variable a∈{1,...,∣M∣}a\in\{1,...,|{\bf M}|\} which selects a memory entry ma{\mathbf{m}}_{a} from the memory buffer M{\bf M}. The varitional posterior distribution q(a∣x)q(a|{\mathbf{x}}) is parameterized as a softmax over a similarity measure between x{\mathbf{x}} and each of the memory entries ma{\mathbf{m}}_{a}:

where Sϕq(x,y)\text{S}^{q}_{\phi}({\mathbf{x}},{\mathbf{y}}) is a learned similarity function described in more detail below.

Given a sample aa from the posterior qϕ(a∣x)q_{\phi}(a|{\mathbf{x}}), retreiving ma{\mathbf{m}}_{a} from MM is a purely deterministic operation. Sampling from q(a∣x)q(a|{\mathbf{x}}) is easy as it amounts to computing its value for each slot in memory and sampling from the resulting categorical distribution. Given aa, we can compute the probability of drawing that address under the prior p(a)p(a). We here use a learned prior p(a)p(a) that shares some parameters with q(a∣x)q(a|{\mathbf{x}}).

Similarity functions: To obtain an efficient implementation for mini-batch training we use the same memory content M{\bf M} for the all training examples in a mini-batch and choose a specific form for the similarity function. We parameterize Sq(m,x)\text{S}^{q}({\mathbf{m}},{\mathbf{x}}) with two MLPs: hϕ\text{h}_{\phi} that embeds the memory content into the matching space and hϕq\text{h}^{q}_{\phi} that does the same to the query x{\mathbf{x}}. The similarity is then computed as the inner product of the embeddings, normalized by the norm of the memory content embedding:

This form allows us to compute the similarities between the embeddings of a mini-batch of ∣B∣|{\cal B}| observations and ∣M∣|{\bf M}| memory entries at the computational cost of O(∣M∣∣B∣∣e∣)O(|{\bf M}||{\cal B}||{\mathbf{e}}|), where ∣e∣|{\mathbf{e}}| is the dimensionality of the embedding. We also experimented with several alternative similarity functions such as the plain inner product (⟨ea,eq⟩\langle{\mathbf{e}}_{a},{\mathbf{e}}^{q}\rangle) and the cosine similarity (\nicefrac⟨ea,eq⟩∣∣ea∣∣⋅∣∣eq∣∣\nicefrac{{\langle{\mathbf{e}}_{a},{\mathbf{e}}^{q}\rangle}}{{||{\mathbf{e}}_{a}||\cdot||{\mathbf{e}}^{q}||}}) and found that they did not outperform the above similarity function. For the unconditioneal prior p(a)p(a), we learn a query point ep∈R∣e∣{\mathbf{e}}^{p}\in{\cal R}^{|e|} to use in similarity function (5) in place of eq{\mathbf{e}}^{q}. We share hϕ\text{h}_{\phi} between p(a)p(a) and q(a∣x)q(a|{\mathbf{x}}). Using a trainable p(a)p(a) allows the model to learn that some memory entries are more useful for generating new targets than others. Control experiments showed that there is only a very small degradation in performance when we assume a flat prior p(a)=\nicefrac1∣M∣p(a)=\nicefrac{{1}}{{|{\bf M}|}}.

For the continuous variable z{\mathbf{z}} we use the methods developed in the context of variational autoencoders (kingma2013vae, ). We use a conditional Gaussian prior p(z∣ma)p(z|{\mathbf{m}}_{a}) and an approximate conditional posterior q(z∣x,ma)q(z|{\mathbf{x}},{\mathbf{m}}_{a}). However, since we have a discrete latent variable aa in the model we can not simply backpropagate gradients through it. Here we show how to use VIMCO (mnih2016variational, ) to estimate the gradients for this model. With VIMCO, we essentially optimize the multi-sample variational bound (bornschein2014reweighted, ; burda2015importance, ; mnih2016variational, ):

Multiple samples from the posterior enable VIMCO to estimate low-variance gradients for those parameters ϕ\phi of the model which influence the non-differentiable discrete variable aa. The corresponding gradient estimates are:

Related work

Attention and external memory are two closely related techniques that have recently become important building blocks for neural models. Attention has been widely used for supervised learning tasks as translation, image classification and image captioning. External memory can be seen as an input or an internal state and attention mechanisms can either be used for selective reading or incremental updating. While most work involving memory and attention has been done in the context supervised learning, here we are interested in using them effectively in the generative setting.

In li2016learning the authors use soft-attention with learned memory contents to augment models to have more parameters in the generative model. External memory as a way of implementing one-shot generalization was introduced in (rezende2016one, ). This was achieved by treating the exemplars conditioned on as memory entries accessed through a soft attention mechanism at each step of the incremental generative process similar to the one in DRAW (gregor2015draw, ). Generative Matching Networks (bartunov2016fast, ) are a similar architecture which uses a single-step VAE generative process instead of an iterative DRAW-like one. In both cases, soft attention is used to access the exemplar memory, with the address weights computed based on a learned similarity function between an observation at the address and a function of the latent state of the generative model.

In contrast to this kind of deterministic soft addressing, we use hard attention, which stochastically picks a single memory entry and thus might be more appropriate in the few-shot setting. As the memory location is stochastic in our model, we perform variational inference over it, which has not been done for memory addressing in a generative model before. A similar approach has however been used for training stochastic attention for image captioning (ba2015learning, ). In the context of memory, hard attention has been used in RLNTM – a version of the Neural Turing Machine modified to use stochastic hard addressing (zaremba2015reinforcement, ). However, RLNTM has been trained using REINFORCE rather than variational inference. A number of architectures for VAEs augmented with mixture priors have been proposed, but they do not use the mixture component indicator variable to index memory and integrate out the variable instead (dilokthanakul2016deep, ; nalisnick2016approximate, ), which prevents them from scaling to a large number of mixing components.

An alternative approach to generative few-shot learning proposed in (harrison2017neuralstatistician, ) uses a hierarchical VAE to model a large number of small related datasets jointly. The statistical structure common to observations in the same dataset are modelled by a continuous latent vector shared among all such observations. Unlike our model, this model is not memory-based and does not use any form of attention. Generative models with memory have also been proposed for sequence modelling in (gemici2017generative, ), using differentiable soft addressing. Our approach to stochastic addressing is sufficiently general to be applicable in this setting as well and it would be interesting how it would perform as a plug-in replacement for soft addressing.

Experiments

We optimize the parameters with Adam (kingma2014adam, ) and report experiments with the best results from learning rates in {1e-4, 3e-4}. We use minibatches of size 32 and KK=4 samples from the approximate posterior q(⋅∣x)q(\cdot|{\mathbf{x}}) to compute the gradients, the KL estimates, and the log-likelihood bounds. We keep the architectures deliberately simple and do not use autoregressive connections or IAF (kingma2016iaf, ) in our models as we are primarily interested in the quantitative and qualitative behaviour of the memory component.

We first perform a series of experiments on the binarized MNIST dataset (LarochelleBinarizedMNIST, ). We use 2 layered en- and decoders with 256 and 128 hidden units with ReLU nonlinearities and a 32 dimensional Gaussian latent variable z{\mathbf{z}}.

Train to recall: To investigate the model’s capability to use its memory to its full extent, we consider the case where it is trained to maximize the likelihood for random data points x{\mathbf{x}} which are present in M{\bf M}. During inference, an optimal model would pick the template ma{\mathbf{m}}_{a} that is equivalent to x{\mathbf{x}} with probability q(a∣x)q(a|{\mathbf{x}})=1. The corresponding prior probability would be p(a)≈\nicefrac1∣M∣p(a)\approx\nicefrac{{1}}{{|{\bf M}|}}. Because there are no further variations that need to be modeled by z{\mathbf{z}}, its posterior q(z∣x,m)q({\mathbf{z}}|{\mathbf{x}},{\mathbf{m}}) can match the prior p(z∣m)p({\mathbf{z}}|{\mathbf{m}}), yielding a KL cost of zero. The model expected log likelihood would be -log⁡∣M∣\log|{\bf M}|, equal to the log-likelihood of an optimal probabilistic lookup table. Figure 2A illustrates that our model converges to the optimal solution. We observed that the time to convergence depends on the size of the memory and with ∣M∣>512|{\bf M}|>512 the model sometimes fails to find the optimal solution. It is noteworthy that the trained model from Figure 2A can handle much larger memory sizes at test time, e.g. achieving NLL ≈log⁡(2048)\approx\log(2048) given 20482048 test set images in memory. This indicates that the matching MLPs for q(a∣x)q(a|{\mathbf{x}}) are sufficiently discriminative.

Learned memory: We train models with ∣M∣∈{64,128,256,512,1024}|{\bf M}|\in\{64,128,256,512,1024\} randomly initialized mixture components (ma∈R256{\mathbf{m}}_{a}\in{\cal R}^{256}). After training, all models converged to an average KL(q(a∣x)∣∣p(a))≈2.5±0.3KL(q(a|{\mathbf{x}})||p(a))\approx 2.5\pm 0.3 nats over both the training and the test set, suggesting that the model identified between e2.2≈9e^{2.2}\approx 9 and e2.8≈16e^{2.8}\approx 16 clusters in the data that are represented by aa. The entropy of p(a)p(a) is significantly higher, indicating that multiple ma{\mathbf{m}}_{a} are used to represent the same data clusters. A manual inspection of the q(a∣x)q(a|{\mathbf{x}}) histograms confirms this interpretation. Although our model overfits slightly more to the training set, we do generally not observe a big difference between our model and the corresponding baseline VAE (a VAE with the same architecture, but without the top level mixture distribution) in terms of the final NLL. This is probably not surprising, because MNIST provides many training examples describing a relatively simple data manifold. Figure 2B shows samples from the model.

2 Omniglot with convolutional MLPs

To apply the model to a more challenging dataset and to use it for generative few-shot learning, we train it on various versions of the Omniglot (lake2015human, ) dataset. For these experiments we use convolutional en- and decoders: The approximate posterior q(z∣m,x)q({\mathbf{z}}|{\mathbf{m}},{\mathbf{x}}) takes the concatenation of x{\mathbf{x}} and m{\mathbf{m}} as input and predicts the mean and variance for the 64 dimensional z{\mathbf{z}}. It consists of 6 convolutional layers with 3×33\times 3 kernels and 48 or 64 feature maps each. Every second layer uses a stride of 2 to get an overall downsampling of 8×88\times 8. The convolutional pyramid is followed by a fully-connected MLP with 1 hidden layer and 2∣z∣2|{\mathbf{z}}| output units. The architecture of p(x∣m,z)p({\mathbf{x}}|{\mathbf{m}},{\mathbf{z}}) uses the same downscaling pyramid to map m{\mathbf{m}} to a ∣z∣|{\mathbf{z}}|-dimensional vector, which is concatenated with z{\mathbf{z}} and upscaled with transposed convolutions to the full image size again. We use skip connections from the downscaling layers of m{\mathbf{m}} to the corresponding upscaling layers to preserve a high bandwidth path from m{\mathbf{m}} to x{\mathbf{x}}. To reduce overfitting, given the relatively small size of the Omniglot dataset, we tie the parameters of the convolutional downscaling layers in q(z∣m)q({\mathbf{z}}|{\mathbf{m}}) and p(x∣m,z)p({\mathbf{x}}|{\mathbf{m}},{\mathbf{z}}). The embedding MLPs for p(a)p(a) and q(a∣x)q(a|{\mathbf{x}}) use the same convolutional architecture and map images x{\mathbf{x}} and memory content ma{\mathbf{m}}_{a} into a 128-dimensional matching space for the similarity calculations. We left their parameters untied because we did not observe any improvement nor degradation of performance when tying them.

With learned memory: We run experiments on the 28×2828\times 28 pixel sized version of Omniglot which was introduced in (burda2015importance, ). The dataset contains 24,345 unlabeled examples in the training, and 8,070 examples in the test set from 1623 different character classes. The goal of this experiment is to show that our architecture can learn to use the top-level memory to model highly multi-modal input data. We run experiments with up to 2048 randomly initialized mixture components and observe that the model makes substantial use of them: The average KL(q(a∣x)∣∣p(a))KL(q(a|{\mathbf{x}})||p(a)) typically approaches log⁡∣M∣\log|{\bf M}|, while KL(q(z∣⋅)∣∣p(z∣⋅))KL(q({\mathbf{z}}|\cdot)||p({\mathbf{z}}|\cdot)) and the overall training-set NLL are significantly lower compared to the corresponding baseline VAE. However big models without regularization tend to overfit heavily (e.g. training-set NLL < 80 nats; testset NLL > 150 nats when using ∣M∣|{\bf M}|=2048). By constraining the model size (∣M∣|{\bf M}|=256, convolutions with 32 feature maps) and adding 3e-4 L2 weight decay to all parameters with the exception of M{\bf M}, we obtain a model with a testset NLL of 103.6 nats (evaluated with K=5000 samples from the posterior), which is about the same as a two-layer IWAE and slightly worse than the best RBMs (103.4 and ≈\approx100 respectively, (burda2015importance, )).

Few-shot learning: The 28×2828\times 28 pixel version (burda2015importance, ) of Omniglot does not contain any alphabet or character-class labels. For few-shot learning we therefore start from the original dataset (lake2015human, ) and scale the 104×104104\times 104 pixel sized examples with 4×44\times 4 max-pooling to 26×2626\times 26 pixels. We here use the 45/5 split introduced in (rezende2016one, ) because we are mostly interested in the quantitative behaviour of the memory component, and not so much in finding optimal regularization hyperparameters to maximize performance on small datasets. For each gradient step, we sample 8 random character-classes from random alphabets. From each character-class we sample 4 examples and use them as targets x{\mathbf{x}} to form a minibatch of size 32. Depending on the experiment, we select a certain number of the remaining examples from the same character classes to populate M{\bf M}. We chose 8 character-classes and 4 examples per class for computational convenience (to obtain reasonable minibatch and memory sizes). In control experiments with 32 character classes per minibatch we obtain almost indistinguishable learning dynamics and results.

To establish meaningful baselines, we train additional models with identical encoder and decoder architectures: 1) A simple, unconditioned VAE. 2) A memory-augmented generative model with soft-attention. Because the soft-attention weights have to depend solely on the variables in the generative model and may not take input directly from the encoder, we have to use z{\mathbf{z}} as the top-level latent variable: p(z),p(x∣z,m(z))p({\mathbf{z}}),p({\mathbf{x}}|{\mathbf{z}},{\mathbf{m}}({\mathbf{z}})) and q(z∣x)q({\mathbf{z}}|{\mathbf{x}}). The overall structure of this model resembles the structure of prior work on memory-augmented generative models (see section 3 and Figure 1A), and is very similar to the one used in bartunov2016fast , for example.

For the unconditioned baseline VAE we obtain a NLL of 90.8, while our memory augmented model reaches up to 68.8 nats. Figure 5 shows the scaling properties of our model when varying the number of conditioning examples at test-time. We observe only minimal degradation compared to a theoretically optimal model when we increase the number of concurrent character classes in memory up to 144, indicating that memory readout works reliably with ∣M∣≥2500|{\bf M}|\geq 2500 items in memory. The soft-attention baseline model reaches up to 73.4 nats when M{\bf M} contains 16 examples from 1 or 2 character-classes, but degrades rapidly with increasing number of confounding classes (see Figure 5A). Figure 3 shows histograms and samples from q(a∣x)q(a|{\mathbf{x}}), visually confirming that our model performs reliable approximate inference over the memory locations.

We also train a model on the Omniglot dataset used in bartunov2016fast . This split provides a relatively small training set. We reduce the number of feature channels and hidden layers in our MLPs and add 3e-4 L2 weight decay to all parameters to reduce overfitting. The model in bartunov2016fast has a clear advantage when many examples from very few character classes are in memory because it was specifically designed to extract joint statistics from memory before applying the soft-attention readout. But like our own soft-attention baseline, it quickly degrades as the number of concurrent classes in memory is increased to 4 (table 1).

Conclusions

In our experiments we generally observe that the proposed model is very well behaved: we never used temperature annealing for the categorical softmax or other tricks to encourage the model to use memory. The interplay between p(a)p(a) and q(a∣x)q(a|x) maintains exploration (high entropy) during the early phase of training and decreases naturally as the sampled ma{\mathbf{m}}_{a} become more informative. The KL divergences for the continuous and discrete latent variables show intuitively interpretable results for all our experiments: On the densely sampled MNIST dataset only a few distinctive mixture components are identified, while on the more disjoint and sparsely sampled Omniglot dataset the model chooses to use many more memory entries and uses the continuous latent variables less. By interpreting memory addressing as a stochastic operation, we gain the ability to apply a variational approximation which helps the model to perform precise memory lookups training through inference. Compared to soft-attention approaches, we loose the ability to naively backprop through read-operations. However, our generative few-shot experiments strongly suggest that this can be a worthwhile trade-off if the memory contains disjoint content that should not be used in interpolation. Our experiments also show that the proposed variational approximation is robust to increasing memory sizes: A model trained with 32 items in memory performed nearly optimally with more than 2500 items in memory at test-time. Beginning with M≥48{\bf M}\geq 48 our implementation using hard-attention becomes noticeably faster in terms of wall-clock time per parameter update than the corresponding soft-attention baseline, even taking account the fact that we use KK=4 posterior samples during training and the soft-attention baseline only requires a single one.

We thank our colleagues at DeepMind and especially Oriol Vinyals and Sergey Bartunov for insightful discussions.

References

Supplement

Approximate posterior inference for Omniglot

Inference with q​(a|𝐱)𝑞conditional𝑎𝐱q(a|{\mathbf{x}}) example:

Inference with q​(a|𝐱)𝑞conditional𝑎𝐱q(a|{\mathbf{x}}) example (incorrect):

Inference with q​(a|𝐱)𝑞conditional𝑎𝐱q(a|{\mathbf{x}}) when target is not in memory:

One-shot samples