KERMIT: Generative Insertion-Based Modeling for Sequences

William Chan, Nikita Kitaev, Kelvin Guu, Mitchell Stern, Jakob Uszkoreit

Introduction

Neural sequence models (Sutskever et al.,, 2014; Cho et al.,, 2014) have been successfully applied to many conditional generation applications, including machine translation (Bahdanau et al.,, 2015; Luong et al.,, 2015), speech recognition (Chan et al.,, 2016; Bahdanau et al.,, 2016), speech synthesis (Oord et al.,, 2016; Wang et al.,, 2017) and image captioning (Vinyals et al.,, 2015; Xu et al.,, 2015). Much of the prior work in this area follows the seq2seq encoder-decoder paradigm, where an encoder builds a representation of an observed sequence xx, and a decoder gives the conditional output distribution p(y∣x)p(y\mid x) according to a predetermined factorization, usually left-to-right.

While effective for straightforward conditional generation, such an approach is inflexible and cannot readily be applied to other inference tasks such as non-left-to-right generation or infilling. In this work, we present a more general approach called Kontextuell Encoder Representations Made by Insertion Transformations, or KERMIT for short. KERMIT is a simple architecture that directly models the joint distribution p(x,y)p(x,y) and its decompositions (such as the marginals p(x)p(x) and p(y)p(y) and the conditionals p(y∣x)p(y\mid x) and p(x∣y)p(x\mid y)) in a unified manner. In contrast with traditional seq2seq models, KERMIT does not rely on a prespecified factorization, but is instead able to condition on whatever information is available and infer what remains.

During training, we present KERMIT with paired data (x,y)(x,y) to learn the joint, and can optionally mix in unpaired data xx or yy to refine the marginals in a semi-supervised setting. At test time, a single KERMIT model can be used for conditional inference in either direction by restricting the output distribution to p(x∣y)p(x\mid y) or p(y∣x)p(y\mid x) as required. We can also generate paired samples from the joint distribution (x,y)∼p(x,y)(x,y)\sim p(x,y), or unpaired samples from the marginals x∼p(x)x\sim p(x) or y∼p(y)y\sim p(y).

KERMIT uses a simple architecture and is easy to implement. It does not have a separate encoder and decoder, nor does it require causality masks. In our implementation, KERMIT consists of a single Transformer decoder stack (Vaswani et al.,, 2017). The model is trained to insert the missing tokens into any partially-complete sequence, as shown in Figure 1. We describe the implementation in more detail in Section 3.

We apply KERMIT to a diverse set of tasks, finding that our unified approach is capable of matching or exceeding the performance of dedicated state-of-the-art systems without the need for problem-specific components. We first apply KERMIT to machine translation, where the inputs and outputs are parallel sentence pairs. Then, like its friends ELMo (Peters et al.,, 2018), BERT (Devlin et al.,, 2019), and ERNIE (Sun et al.,, 2019), we can also use KERMIT for self-supervised representation learning for use in downstream NLP tasks. Finally, we apply KERMIT to a zero-shot cloze question-answering task demonstrating the infilling capabilities of the model. Table 1 summarizes our results on all three tasks compared to other highly tuned models: Transformer, BERT, GPT and GPT-2.

Background

In this section, we define some notation and give a brief review of existing sequence models, including autoregressive left-to-right models (Sutskever et al.,, 2014; Cho et al.,, 2014) and masked language models (Devlin et al.,, 2019).

Let X\mathcal{X} and Y\mathcal{Y} be the set of all input and output sequences, respectively. In a standard sequence-to-sequence task, we are presented with training data consisting of sequence pairs (x,y)∈X×Y(x,y)\in\mathcal{X}\times\mathcal{Y}, e.g. parallel translations, and we aim to learn the conditional distribution p(y∣x)p(y\mid x). Traditional autoregressive models (Sutskever et al.,, 2014; Cho et al.,, 2014) use a left-to-right factorization, decomposing the distribution as a chain of predictions conditioning on the input xx and prefixes y<ty_{<t}:

This structure is also used for unconditional sequence tasks such as language modeling where the goal is to learn an unconditional output distribution on its own. A left-to-right factorization is convenient because it allows for exact log-likelihood computation, thereby permitting efficient maximum likelihood estimation. It also leads to simple approximate inference algorithms such as greedy decoding

or beam search over sets of multiple hypotheses.

However, there are some drawbacks to the autoregressive approach. First, in the case of conditional generation, it cannot handle situations where the input xx is only partially observed. Second, since it utilizes a fixed left-to-right factorization, it cannot be used for other inference tasks like infilling where generation is not monotonic. Moreover, standard inference algorithms require nn generation steps to generate nn tokens, which could be a bottleneck in end-use applications.

2 Masked Language Models

Masked Language Models (MLMs) (Devlin et al.,, 2019) comprise another class of models targeting the unconditional setting. For MLMs, a partial canvas xs⊆xx_{s}\subseteq x is observed where some of the tokens in xx have been masked out, and the objective is to recover xx from xsx_{s}. For example, for a ground truth canvas x∗=(A,B,C,D,E)x^{*}=(A,B,C,D,E) and a partial canvas xs∗=(A,_,C,D,_)x^{*}_{s}=(A,\_,C,D,\_), the model should learn to replace the second blank with BB and the last blank with EE. The model outputs an independent prediction at each position, and its objective is to maximize p(x∣xs)p(x\mid x_{s}).

Because the exact locations of the slots are known in xsx_{s}, the model does not need to predict where the missing items are located, but only what they should be. Consequently, the model is not immediately suitable for generation, as the canvas size needs to be fixed during inference and cannot change over time (i.e., ∣xs∣=∣x∣|x_{s}|=|x|). MLMs have been successfully applied in self-supervised representation learning settings, leading to strong results on downstream language tasks (Devlin et al.,, 2019).

KERMIT

In this section we propose KERMIT, a novel insertion-based generative model. Unlike the prior work mentioned in Section 2, KERMIT does not have the rigid construction of modeling the target sequence given some fully observed source sequence, nor does it assume a left-to-right factorization (and generation order) of the output sequence. To motivate and arrive at our model, we formalize then extend a recent insertion-based conditional modeling framework proposed by Stern et al., (2019).

We begin with the unconditional setting. In order to model sequences without requiring a fixed factorization or imposing constraints on the order of generation, we make use of a framework in which sequences are constructed via insertion operations. Given a sequence x=(x1,…,xn)x=(x_{1},\dots,x_{n}) and a generation order zz represented as a permutation of the indices {1,…,n}\{1,\dots,n\}, we define the corresponding sequence ((c1z,l1z),…,(cnz,lnz))((c^{z}_{1},l^{z}_{1}),\dots,(c^{z}_{n},l^{z}_{n})) of insertion operations which produces xx according to order zz. Here, ciz∈Cc^{z}_{i}\in\mathcal{C} is an element of the vocabulary and 1≤liz≤i1\leq l^{z}_{i}\leq i is an insertion location relative to the current hypothesis. For example, if constructing the sequence (A,B,C)(A,B,C) as ()→(C)→(A,C)→(A,B,C)()\to(C)\to(A,C)\to(A,B,C), we would have z=(3,1,2)z=(3,1,2) with (c1z,l1z)=(C,1),(c2z,l2z)=(A,1),(c3z,l3z)=(B,2)(c^{z}_{1},l^{z}_{1})=(C,1),(c^{z}_{2},l^{z}_{2})=(A,1),(c^{z}_{3},l^{z}_{3})=(B,2).

Next let (x1z,i,…,xiz,i)(x^{z,i}_{1},\dots,x^{z,i}_{i}) denote the subsequence of xx corresponding to the (ordered) extraction of the elements at indices {z1,…,zi}\{z_{1},\dots,z_{i}\}. This is the partial output at iteration ii. Note that this will be the same for all permutations zz with the same unordered set of indices in the first ii positions. For the example above for instance, we have (x1z,2,x2z,2)=(A,C)(x^{z,2}_{1},x^{z,2}_{2})=(A,C).

Armed with these definitions, we can now write out p(x)p(x) as a marginalization over all possible orders z∈Snz\in S_{n} for sequence length nn, where SnS_{n} denotes the set of all permutations on nn elements:

where the last line encodes the Markov assumption that the order of insertions leading to a given canvas is not important, just the result. Typically we will use a uniform prior over permutations for p(z)p(z), though other options are available, such as the balanced binary tree prior described by Stern et al., (2019).

Although exact computation of the log-likelihood is intractable due to the marginalization over the generation order zz, we can lower bound the log-likelihood using Jensen’s inequality via

where the simplification in the last line follows from the fact that ∑zi+1:np(zi+1:n∣z1:i)=1\sum_{z_{i+1:n}}p(z_{i+1:n}\mid z_{1:i})=1.

From here, we can multiply and divide the outer sum by nn to turn it into a mean, then arrive at the following simple sampling procedure to compute an unbiased estimate of our lower bound L(x)\mathcal{L}(x) on the log-likelihood for a single example:

Sample a partial permutation z1:i−1∼p(z1:i−1)z_{1:i-1}\sim p(z_{1:i-1}) for the first i−1i-1 insertions.

Compute a weighted sum over the next-step losses log⁡p((ciz,liz)∣x1:i−1z,i−1)\log p((c^{z}_{i},l^{z}_{i})\mid x^{z,i-1}_{1:i-1}) scaled by the weighting distribution p(zi∣z1:i−1)p(z_{i}\mid z_{1:i-1}) and the sequence length nn.

2 Inference

Using this model, inference can be autoregressive via greedy decoding

or partially autoregressive via parallel decoding

In the case of parallel decoding, we perform simultaneous insertions at all non-finished slots. If we use a balanced binary tree prior for p(z)p(z) (Stern et al.,, 2019), we can even achieve an empirical runtime of ≈log⁡2n\approx\log_{2}n iterations to generate nn tokens. One key advantage of insertion-based models over MLMs is that the output canvas can dynamically grow in size, meaning the length does not need to be chosen before the start of generation.

3 Pairs of Sequences

By keeping our architecture order-agnostic and marginalizing over all possible orders in our training objective, KERMIT is able to learn the joint distribution and all its decompositions, including the marginals p(x)p(x) and p(y)p(y) and conditionals p(y∣x)p(y\mid x) and p(x∣y)p(x\mid y). We can also perform targeted training. More explicitly, if the model is provided with a canvas that fully contains xx or yy, then it will learn a conditional distribution. If the model is provided with an example where xx or yy is empty, then it will learn the opposing marginal distribution.

4 Model

We implement KERMIT as a single Transformer decoder stack (Vaswani et al.,, 2017), without any form of causal masking. The full self-attention mechanism allows the model to capture any relationships between the input canvas and the predicted insertion operations with a constant number of operations. We follow Stern et al., (2019) and model the (content, location) distribution p(c,l)p(c,l) as a factorized distribution p(c,l)=p(c∣l)p(l)p(c,l)=p(c\mid l)p(l), where p(c∣l)p(c\mid l) is the standard Transformer softmax over the vocabulary, and a p(l)p(l) is a softmax over the locations. Figure 2 visualizes the differences between a standard Transformer (Vaswani et al.,, 2017), BERT (Devlin et al.,, 2019), Insertion Transformer (Stern et al.,, 2019) and KERMIT.

Experiments

We perform experiments with KERMIT on the tasks of machine translation, self-supervised representation learning, and zero-shot cloze question answering.

We first apply KERMIT on the competitive WMT 2014 English ↔\leftrightarrow German translation task. We follow the hyperparameter settings of the base Transformer (Vaswani et al.,, 2018). However, since KERMIT does not have an encoder, we simply double the decoder width. We perform no additional hyperparameter tuning. We also follow prior work (Gu et al.,, 2018; Stern et al.,, 2018, 2019; Lee et al.,, 2018) in using distillation (Hinton et al.,, 2015; Kim and Rush,, 2016) to train our models. We follow Stern et al., (2019) in using a balanced binary tree loss, and we similarly observe an empirically logarithmic number of generation steps in sequence length when using parallel decoding. However, unlike Stern et al., (2019) we did not need to tune an EOS penalty, but simply set it to zero for all experiments.

We train several different KERMIT models for translation. First we train two unidirectional models, where the model observes a full source sentence (i.e., English or German) and is asked to generate the corresponding target sentence (i.e., German or English). These separately learn the conditional distributions p(y∣x)p(y\mid x) and p(x∣y)p(x\mid y), mimicking the traditional conditional generation setup. On the WMT 2014 test set, we achieve 27.8/30.7 BLEU with this approach, roughly matching our base Transformer baseline of 27.8/31.2 BLEU. We also train a bidirectional model on the union of the two unidirectional training sets, yielding a single model that captures both conditional distributions p(y∣x)p(y\mid x) and p(x∣y)p(x\mid y). We do not change any hyperparameters when training this model (i.e., we do not increase model capacity). The combined approach obtains 27.2/27.6 BLEU, nearly matching the baseline for English →\to German but falling slightly behind in the reverse direction.

We also train a full joint model that captures the full joint distribution p(x,y)p(x,y) and factorizations thereof. Like the bidirectional model, the joint model can translate in either direction, but it can additionally be used for sampling or completing partial inputs. We use the same hyperparameter set as before. Since the model is now faced with a much more challenging task, it does slightly worse when limited to the same model size, but still reaches a respectable 25.6/27.4 BLEU. Unlike the previous models, however, we can incorporate monolingual data into the joint model’s training setup to supplement its knowledge of the marginals p(x)p(x) and p(y)p(y). We accordingly train a joint model with all our paired data and 1M additional samples of English and German monolingual data randomly selected from the WMT 2014 monolingual corpus. Without altering model capacity, we find that refining the marginals gives us a 1.2 BLEU improvement on German →\rightarrow English. Finally, we take the model which was trained on the full joint distribution with marginal refinement, and further finetune it on both the unidirectional and bidirectional settings. We find a small improvement in BLEU over the original models in both settings.

Table 2 summarizes our results. We emphasize that virtually all of our models outperform prior non-fully-autoregressive approaches in terms of BLEU. We also note that the observed number of iterations required to generate nn tokens is roughly log⁡2n\log_{2}n due to the use of a balanced binary tree loss and parallel decoding, which is substantially lower than autoregressive models which require nn steps. Some examples of parallel decodes are shown in Figure 3. Our models require an average of 5.5-6.5 decoding iterations for the sentences in the test set, outperforming the constant-time models of Lee et al., (2018) which require 10 iterations in both BLEU and empirical decoding complexity.

We also draw samples from the model to highlight its infilling and generation capabilities. Figure 4 captures some examples. We first show unconditional sampling of an (English, German) sentence pair. We also take a translation example from the newstest2013 dev set and split it in half, sampling completions after seeding the English side with the first half and the German side with the second half. We find the model is capable of generating a very diverse set of coherent samples.

2 Representation Learning

Like its close friend BERT (Devlin et al.,, 2019), KERMIT can also be used for self-supervised representation learning and applied to various language understanding tasks. We follow the same training procedure and hyperparameter setup as BERT\textsclarge{}_{\textsc{large}}. However, instead of masking 15% of the tokens and replacing them with blank tokens like in BERT (Devlin et al.,, 2019), KERMIT simply drops them out completely from the sequence.

Prior to BERT, the best representation learning approach was to use a language model such as GPT (Radford et al.,, 2018). BERT outperforms GPT in large part because of its deeply bi-directional architecture, but in the process BERT sacrifices the ability to perform straightforward generation. While we find KERMIT to perform slightly behind BERT, KERMIT maintains the ability to generate text while obtaining results that are much closer to BERT rather than GPT. The GLUE benchmark (Wang et al.,, 2019) results are summarized in Table 3.

3 Zero-Shot Cloze Question Answering

Finally, we also investigate the infilling abilities of KERMIT and related approaches by evaluating their performance on zero-shot cloze question answering. In particular, we aim to understand how effective these models are for fill-in-the-blank-style question answering after being trained only on language modeling data without any task-specific fine-tuning.

For this experiment, we use the human-annotated QA2D dataset assembled by Demszky et al., (2018), which consists of examples from the SQuAD dataset (Rajpurkar et al.,, 2016) in which the answer has been extended from a single phrase into a full declarative sentence. These can be transformed into cloze instances by removing the answer phrase from the declarative output. For example, given the question “When was Madonna born?” and the answer “August 16, 1958”, the full declarative answer would be “Madonna was born on August 16, 1958.”, and the associated cloze instance would be “Madonna was born on .”

We take the KERMIT model trained from Section 4.2 and two powerful language models (BERT (Devlin et al.,, 2019) and the largest public version of GPT-2The 345M parameter “medium size” model. (Radford et al.,, 2019)), and evaluate their ability to fill in the blank of each cloze instance, each without specifically being trained on data of this form. We employ different decoding strategies as required for each model, detailed below.

We split the passage in half and present KERMIT with examples of the form

where cloze(left) and cloze(right) are the portions of the declarative answer before and after the gap. Since KERMIT can natively perform insertions, we simply perform a parallel decode constrained to take place within the gap and extract the output as our answer.

We split the passage in half and present BERT with examples of the form

Here we include explicit [MASK] tokens, running separate decodes with n=1,2,…n=1,2,\dots up to 44 or the oracle answer length, whichever is greater. We then choose the one with the highest score under the model and extract the outputs at the masked positions as the answer. Each decode consists of a beam search in which one [MASK] is filled at a time. For each element on the beam, we choose the remaining [MASK] position with the highest confidence (lowest entropy) as the next position to fill. We found that this beam-search did substantially better than left-to-right decoding or parallel decoding.

For GPT-2, a left-to-right language model, we cannot directly condition on both the left and right context. Instead, we first present the model with the prefix

and sample continuations of varying lengths. For each continuation, we then append cloze(right) and compute the score of the full sequence under the model. We select the best-scoring sequence and extract the portion in the gap as the answer. To efficiently obtain continuations of varying lengths, we generate 20 extended continuations from the model, then treat all prefixes of those continuations as candidate values to go in the gap.

We evaluate on 50,000 cloze-formulated questions from SQuAD, using the standard SQuAD evaluation script to compute accuracy in terms of exact match and token-level F1. Results are presented in Table 4. KERMIT performs significantly better on this zero-shot cloze task than the other two approaches thanks to its infilling capabilities learned through its insertion-oriented objective, achieving 30.3 F1 and 20.9% exact match. BERT’s performance falls short of KERMIT, as it often prefers shorter completions since it is not required to handle length modeling during training. GPT-2 lags further behind the others due to its inability to condition on the context on both sides of the gap during inference. Even when the oracle length (i.e., the ground-truth length of the answer) is provided to BERT and GPT-2, KERMIT still substantially outperforms all other models.

Conclusion

In this paper, we present KERMIT, an insertion-based framework for sequences that can model the joint data distribution and its decompositions (i.e., marginals and conditionals). KERMIT can generate text in an arbitrary order – including bidirectional machine translation and cloze-style infilling – and empirically can generate sequences in logarithmic time. It uses a simple neural architecture that can additionally produce contextualized vector representations of words and sentences. We find KERMIT is capable of matching or exceeding state-of-the-art performance on three diverse tasks: machine translation, representation learning, and zero shot cloze question answering.

We give thanks to Samy Bengio, Zhifeng Chen, Jamie Kiros, Luheng He, Geoffrey Hinton, Quoc Le, Lala Li, Mohammad Norouzi, Yu Zhang, and the Google Brain team for useful discussions and technical assistance. Special thanks to Jamie Kiros for brainstorming the name KERMIT.

References