Latent Alignment and Variational Attention

Yuntian Deng, Yoon Kim, Justin Chiu, Demi Guo, Alexander M. Rush

Introduction

Attention networks have quickly become the foundation for state-of-the-art models in natural language understanding, question answering, speech recognition, image captioning, and more . Alongside components such as residual blocks and long-short term memory networks, soft attention provides a rich neural network building block for controlling gradient flow and encoding inductive biases. However, more so than these other components, which are often treated as black-boxes, researchers use intermediate attention decisions directly as a tool for model interpretability or as a factor in final predictions . From this perspective, attention plays the role of a latent alignment variable . An alternative approach, hard attention , makes this connection explicit by introducing a latent variable for alignment and then optimizing a bound on the log marginal likelihood using policy gradients. This approach generally performs worse (aside from a few exceptions such as ) and is used less frequently than its soft counterpart.

Still the latent alignment approach remains appealing for several reasons: (a) latent variables facilitate reasoning about dependencies in a probabilistically principled way, e.g. allowing composition with other models, (b) posterior inference provides a better basis for model analysis and partial predictions than strictly feed-forward models, which have been shown to underperform on alignment in machine translation , and finally (c) directly maximizing marginal likelihood may lead to better results.

The aim of this work is to quantify the issues with attention and propose alternatives based on recent developments in variational inference. While the connection between variational inference and hard attention has been noted in the literature , the space of possible bounds and optimization methods has not been fully explored and is growing quickly. These tools allow us to better quantify whether the general underperformance of hard attention models is due to modeling issues (i.e. soft attention imbues a better inductive bias) or optimization issues.

Our main contribution is a variational attention approach that can effectively fit latent alignments while remaining tractable to train. We consider two variants of variational attention: categorical and relaxed. The categorical method is fit with amortized variational inference using a learned inference network and policy gradient with a soft attention variance reduction baseline. With an appropriate inference network (which conditions on the entire source/target), it can be used at training time as a drop-in replacement for hard attention. The relaxed version assumes that the alignment is sampled from a Dirichlet distribution and hence allows attention over multiple source elements.

Experiments describe how to implement this approach for two major attention-based models: neural machine translation and visual question answering (Figure 1 gives an overview of our approach for machine translation). We first show that maximizing exact marginal likelihood can increase performance over soft attention. We further show that with variational (categorical) attention, alignment variables significantly surpass both soft and hard attention results without requiring much more difficult training. We further explore the impact of posterior inference on alignment decisions, and how latent variable models might be employed. Our code is available at https://github.com/harvardnlp/var-attn/.

Related Work Latent alignment has long been a core problem in NLP, starting with the seminal IBM models , HMM-based alignment models , and a fast log-linear reparameterization of the IBM 2 model . Neural soft attention models were originally introduced as an alternative approach for neural machine translation , and have subsequently been successful on a wide range of tasks (see for a review of applications). Recent work has combined neural attention with traditional alignment and induced structure/sparsity , which can be combined with the variational approaches outlined in this paper.

In contrast to soft attention models, hard attention approaches use a single sample at training time instead of a distribution. These models have proven much more difficult to train, and existing works typically treat hard attention as a black-box reinforcement learning problem with log-likelihood as the reward . Two notable exceptions are : both utilize amortized variational inference to learn a sampling distribution which is used obtain importance-sampled estimates of the log marginal likelihood . Our method uses uses different estimators and targets the single sample approach for efficiency, allowing the method to be employed for NMT and VQA applications.

There has also been significant work in using variational autoencoders for language and translation application. Of particular interest are those that augment an RNN with latent variables (typically Gaussian) at each time step and those that incorporate latent variables into sequence-to-sequence models . Our work differs by modeling an explicit model component (alignment) as a latent variable instead of auxiliary latent variables (e.g. topics). The term "variational attention" has been used to refer to a different component the output from attention (commonly called the context vector) as a latent variable , or to model both the memory and the alignment as a latent variable . Finally, there is some parallel work which also performs exact/approximate marginalization over latent alignments for sequence-to-sequence learning.

Background: Latent Alignment and Neural Attention

We begin by introducing notation for latent alignment, and then show how it relates to neural attention. For clarity, we are careful to use alignment to refer to this probabilistic model (Section 2.1), and soft and hard attention to refer to two particular inference approaches used in the literature to estimate alignment models (Section 2.2).

Directly maximizing this log marginal likelihood in the presence of the latent variable zz is often difficult due to the expectation (though tractable in certain cases).

This computation requires a factor of O(T)O(T) additional runtime, and introduces a major computational factor into already expensive deep learning models.Although not our main focus, explicit marginalization is sometimes tractable with efficient matrix operations on modern hardware, and we compare the variational approach to explicit enumeration in the experiments. In some cases it is also possible to efficiently perform exact marginalization with dynamic programming if one imposes additional constraints (e.g. monotonicity) on the alignment distribution .

2 Attention Models: Soft and Hard

When training deep learning models with gradient methods, it can be difficult to use latent alignment directly. As such, two alignment-like approaches are popular: soft attention replaces the probabilistic model with a deterministic soft function and hard attention trains a latent alignment model by maximizing a lower bound on the log marginal likelihood (obtained from Jensen’s inequality) with policy gradient-style training. We briefly describe how these methods fit into this notation.

Soft attention networks use an altered model shown in Figure 2b. Instead of using a latent variable, they employ a deterministic network to compute an expectation over the alignment variable. We can write this model using the same functions ff and aa from above,

The proof is given in Appendix A.It is also possible to study the gap in finer detail by considering distributions over the inputs of ff that have high probability under approximately linear regions of ff, leading to the notion of approximately expectation-linear functions, which was originally proposed and studied in the context of dropout . Empirically the soft approximation works remarkably well, and often moves towards a sharper distribution with training. Alignment distributions learned this way often correlate with human intuition (e.g. word alignment in machine translation) .Another way of viewing soft attention is as simply a non-probabilistic learned function. While it is possible that such models encode better inductive biases, our experiments show that when properly optimized, latent alignment attention with explicit latent variables do outperform soft attention.

Hard Attention

Variational Attention for Latent Alignment Models

Amortized variational inference (AVI, closely related to variational auto-encoders) is a class of methods to efficiently approximate latent variable inference, using learned inference networks. In this section we explore this technique for deep latent alignment models, and propose methods for variational attention that combine the benefits of soft and hard attention.

With the right choice of optimization strategy and inference network this form of variational attention can provide a general method for learning latent alignment models. In the rest of this section, we consider strategies for accurately and efficiently computing this objective; in the next section, we describe instantiations of encenc for specific domains.

Many recent methods target this issue, including neural estimates of baselines , Rao-Blackwellization , reparameterizable relaxations , and a mix of various techniques . We found that an approach using REINFORCE along with a specialized baseline was effective. However, note that REINFORCE is only one of the inference choices we can select, and as we will show later, alternative approaches such as reparameterizable relaxations work as well. Formally, we first apply the likelihood-ratio trick to obtain an expression for the gradient with respect to the inference network parameters ϕ\phi,

Effectively this weights gradients to qq based on the ratio of the inference network alignment approach to a soft attention baseline. Notably the expectation in the soft attention is over pp (and not over qq), and therefore the baseline is constant with respect to ϕ\phi. Note that a similar baseline can also be used for hard attention, and we apply it to both variational/hard attention models in our experiments.

Algorithm 2: Relaxed Alignments

Next consider treating both D\mathcal{D} and Q\cal Q as Dirichlets, where zz represents a mixture of indices. This model is in some sense closer to the soft attention formulation which assigns mass to multiple indices, though fundamentally different in that we still formally treat alignment as a latent variable. Again the aim is to find a low variance gradient estimator. Instead of using REINFORCE, certain continuous distributions allow the use reparameterization , where sampling z∼q(z)z\sim q(z) can be done by first sampling from a simple unparameterized distribution U\mathcal{U}, and then applying a transformation gϕ(⋅)g_{\phi}(\cdot), yielding an unbiased estimator,

The Dirichlet distribution is not directly reparameterizable. While transforming the standard uniform distribution with the inverse CDF of Dirichlet would result in a Dirichlet distribution, the inverse CDF does not have an analytical solution. However, we can use rejection based sampling to get a sample, and employ implicit differentiation to estimate the gradient of the CDF .

Models and Methods

We experiment with variational attention in two different domains where attention-based models are essential and widely-used: neural machine translation and visual question answering.

For variational attention, the inference network applies a bidirectional LSTM over the source and the target to obtain the hidden states x1,…,xTx_{1},\dots,x_{T} and h1,…,hSh_{1},\dots,h_{S}, and produces the alignment scores at the jj-th time step via a bilinear map, si(j)=exp⁡(hj⊤Uxi)s_{i}^{(j)}=\exp(h_{j}^{\top}\mathbf{U}x_{i}). For the categorical case, the scores are normalized, q(zi(j)=1)∝si(j)q(z^{(j)}_{i}=1)\propto s_{i}^{(j)}; in the relaxed case the parameters of the Dirichlet are αi(j)=si(j)\alpha_{i}^{(j)}=s_{i}^{(j)}. Note, the inference network sees the entire target (through bidirectional LSTMs). The word embeddings are shared between the generative/inference networks, but other parameters are separate.

Visual Question Answering

where ⊙\odot is the element-wise product. This parameterization worked better than alternatives. We did not experiment with the relaxed case in VQA, as the object bounding boxes already give us the ability to attend to larger portions of the image.

Inference Alternatives

For categorical alignments we described maximizing a particular variational lower bound with REINFORCE. Note that other alternatives exist, and we briefly discuss them here: 1) instead of the single-sample variational bound we can use a multiple-sample importance sampling based approach such as Reweighted Wake-Sleep (RWS) or VIMCO ; 2) instead of REINFORCE we can approximate sampling from the discrete categorical distribution with Gumbel-Softmax ; 3) instead of using an inference network we can directly apply Stochastic Variational Inference (SVI) to learn the local variational parameters in the posterior.

Predictive Inference

Experiments

For NMT we mainly use the IWSLT dataset . This dataset is relatively small, but has become a standard benchmark for experimental NMT models. We follow the same preprocessing as in with the same Byte Pair Encoding vocabulary of 14k tokens . To show that variational attention scales to large datasets, we also experiment on the WMT 2017 English-German dataset , following the preprocessing in except that we use newstest2017 as our test set. For VQA, we use the VQA 2.0 dataset. As we are interested in intrinsic evaluation (i.e. log-likelihood) in addition to the standard VQA metric, we randomly select half of the standard validation set as the test set (since we need access to the actual labels). VQA eval metric is defined as min⁡{# humans that said answer 3,1}\min\{\frac{\#\text{ humans that said answer }}{3},1\}. Also note that since there are sometimes multiple answers for a given question, in such cases we sample (where the sampling probability is proportional to the number of humans that said the answer) to get a single label. (Therefore the numbers provided are not strictly comparable to existing work.) While the preprocessing is the same as , our numbers are worse than previously reported as we do not apply any of the commonly-utilized techniques to improve performance on VQA such as data augmentation and label smoothing.

Experiments vary three components of the systems: (a) training objective and model, (b) training approximations, comparing enumeration or sampling,Note that enumeration does not imply exact if we are enumerating an expectation on a lower bound. (c) test inference. All neural models have the same architecture and the exact same number of parameters θ\theta (the inference network parameters ϕ\phi vary, but are not used at test). When training hard and variational attention with sampling both use the same baseline, i.e the output from soft attention. The full architectures/hyperparameters for both NMT and VQA are given in Appendix B.

Results and Discussion

Table 1 shows the main results. We first note that hard attention underperforms soft attention, even when its expectation is enumerated. This indicates that Jensen’s inequality alone is a poor bound. On the other hand, on both experiments, exact marginal likelihood outperforms soft attention, indicating that when possible it is better to have latent alignments.

Table 2 (left) considers test inference for variational attention, comparing enumeration to KK-max with K=5K=5. For all methods exact enumeration is better, however KK-max is a reasonable approximation. Table 2 (right) shows the PPL of different models as we increase KK. Good performance requires K>1K>1, but we only get marginal benefits for K>5K>5. Finally, we observe that it is possible to train with soft attention and test using KK-Max with a small performance drop (Soft KMax in Table 2 (right)). This possibly indicates that soft attention models are approximating latent alignment models. On the other hand, training with latent alignments and testing with soft attention performed badly.

Table 3 (lower right) looks at the entropy of the prior distribution learned by the different models. Note that hard attention has very low entropy (high certainty) whereas soft attention is quite high. The variational attention model falls in between. Figure 3 (left) illustrates the difference in practice.

Table 3 (upper right) compares inference alternatives for variational attention. RWS reaches a comparable performance as REINFORCE, but at a higher memory cost as it requires multiple samples. Gumbel-Softmax reaches nearly the same performance and seems like a viable alternative; although we found its performance is sensitive to its temperature parameter. We also trained a non-amortized SVI model, but found that at similar runtime it was not able to produce satisfactory results, likely due to insufficient updates of the local variational parameters. A hybrid method such as semi-amortized inference might be a potential future direction worth exploring.

Despite extensive experiments, we found that variational relaxed attention performed worse than other methods. In particular we found that when training with a Dirichlet KL, it is hard to reach low-entropy regions of the simplex, and the attentions are more uniform than either soft or variational categorical attention. Table 3 (lower right) quantifies this issue. We experimented with other distributions such as Logistic-Normal and Gumbel-Softmax but neither fixed this issue. Others have also noted difficulty in training Dirichlet models with amortized inference .

Potential Limitations

While this technique is a promising alternative to soft attention, there are some practical limitations: (a) Variational/hard attention needs a good baseline estimator in the form of soft attention. We found this to be a necessary component for adequately training the system. This may prevent this technique from working when TT is intractably large and soft attention is not an option. (b) For some applications, the model relies heavily on having a good posterior estimator. In VQA we had to utilize domain structure for the inference network construction. (c) Recent models such as the Transformer , utilize many repeated attention models. For instance the current best translation models have the equivalent of 150 different attention queries per word translated. It is unclear if this approach can be used at that scale as predictive inference becomes combinatorial.

Conclusion

Attention methods are ubiquitous tool for areas like natural language processing; however they are difficult to use as latent variable models. This work explores alternative approaches to latent alignment, through variational attention with promising result. Future work will experiment with scaling the method on larger-scale tasks and in more complex models, such as multi-hop attention models, transformer models, and structured models, as well as utilizing these latent variables for interpretability and as a way to incorporate prior knowledge.

Acknowledgements

We are grateful to Sam Wiseman and Rachit Singh for insightful comments and discussion, as well as Christian Puhrsch for help with translations. This project was supported by a Facebook Research Award (Low Resource NMT). YK is supported by a Google AI PhD Fellowship. YD is supported by a Bloomberg Research Award. AMR gratefully acknowledges the support of NSF CCF-1704834 and an Amazon AWS Research award.

References

Appendix A: Proof of Proposition 1

where c=max⁡{∣λmax⁡∣,∣λmin⁡∣}c=\max\{|\lambda_{\max}|,|\lambda_{\min}|\} is the largest absolute eigenvalue of Hgx,y^(z^)H_{g_{x,\hat{y}}}(\hat{z}). (Here λmax⁡\lambda_{\max} and λmin⁡\lambda_{\min} are maximum/minimum eigenvalues of HgX,q(z^)H_{g_{X,q}}(\hat{z})). Note that cc is also equal to the spectral norm ∥HgX,q(z^)∥2\|H_{g_{X,q}}(\hat{z})\|_{2} since the Hessian is symmetric.

Here the first inequality follows due to the convexity of the absolute value function and the last inequality follows since

Appendix B: Experimental Setup

For data processing we closely follow the setup in , which uses Byte Pair Encoding over the combined source/target training set to obtain a vocabulary size of 14,000 tokens. However, different from which uses maximum sequence length of 175, for faster training we only train on sequences of length up to 125.

The encoder is a two-layer bi-directional LSTM with 512 units in each direction, and the decoder as a two-layer LSTM with with 768 units. For the decoder, the convex combination of source hidden states at each time step from the attention distribution is used as additional input at the next time step. Word embedding is 512-dimensional.

The inference network consists of two bi-directional LSTMs (also two-layer and 512-dimensional each) which is run over the source/target to obtain the hidden states at each time step. These hidden states are combined using bilinear attention to produce the variational parameters. (In contrast the generative model uses MLP attention from , though we saw little difference between the two parameterizations). Only the word embedding is shared between the inference network and the generative model.

Other training details include: batch size of 6, dropout rate of 0.3, parameter initialization over a uniform distribution U[−0.1,0.1]\mathcal{U}[-0.1,0.1], gradient norm clipping at 5, and training for 30 epochs with Adam (learning rate = 0.0003, β1=\beta_{1}= 0.9, β2=\beta_{2}= 0.999) with a learning rate decay schedule which starts halving the learning rate if validation perplexity does not improve. Most models converged well before 30 epochs.

For decoding we use beam search with beam size 10 and length penalty α=1\alpha=1, from . The length penalty added about 0.5 BLEU points across all the models.

Visual Question Answering

The model first obtains object features by mean-pooling the pretrained ResNet-101 features (which are 2048-dimensional) over object regions given by Faster R-CNN .The ResNet features are kept fixed and not fine-tuned during training. We fix the maximum number of possible regions to be 36. For the question embedding we use a one-layer LSTM with 1024 units over word embeddings. The word embeddings are 300-dimensional and initialized with GloVe . The generative model produces a distribution over the possible objects via applying MLP attention, i.e.

The selected image region is concatenated with the question embedding and fed to a one-layer MLP with ReLU non-linearity and 1024 hidden units.

Other training details include: batch size of 512, dropout rate of 0.5 on the penultimate layer (i.e. before affine transformation into answer vocabulary), and training for 50 epochs with with Adam (learning rate = 0.0005, β1=\beta_{1}= 0.9, β2=\beta_{2}= 0.999) .

In cases where there is more than one answer for a given question/image pair, we randomly sample the answer, where the sampling probability is proportional to the number of humans who gave the answer.

Appendix C: Additional Visualizations