Self-conditioned Embedding Diffusion for Text Generation
Robin Strudel, Corentin Tallec, Florent Altché, Yilun Du, Yaroslav Ganin, Arthur Mensch, Will Grathwohl, Nikolay Savinov, Sander Dieleman, Laurent Sifre, Rémi Leblond
Introduction
Continuous diffusion models (Sohl-Dickstein et al., 2015) have taken the world of image generation by storm, advancing the state of the art further than ever before (Rombach et al., 2021; Ramesh et al., 2022). Can the same framework encounter as much success on the text modality? Diffusion for language is indeed an attractive prospect. Compared to autoregressive (AR) models (Bengio et al., 2000; Sutskever et al., 2011; Austin et al., 2021; Hoffmann et al., 2022), diffusion models can predict all tokens in a sequence at once. This allows for bidirectional, rather than causal attention—increasing interactions between tokens, potentially leading to more coherent samples. Diffusion models can make a better usage of hardware accelerators during inference than AR models, since computations are parallelizable over the sequence axis.
Yet AR models remain the mainstream approach for modelling text. A major obstacle to text diffusion is that diffusion processes typically operate in continuous space. While this naturally handle images, text is inherently discrete. Consequently, most previous attempts to apply diffusion to text have focused on discrete diffusion-like approaches. These methods do not benefit from the refinements made to continuous diffusion in the image domain. Crucially, they cannot make use of guidance (Dhariwal & Nichol, 2021), which drastically improves diffusion models sample quality.
We address this gap by making a simple observation: language models operate mostly in continuous space, with discrete tokens only as inputs and outputs. A natural idea is then to conduct diffusion directly in a continuous token embedding space. For simplicity, we use a fixed embedding space, either random or stemming from a trained language model. Combined with the “self-conditioning” (Chen et al., 2022) refinement, this forms the basis of the method we propose, Self-conditioned Embedding Diffusion (Sed).
Sed models rival mainstream AR models in both conditional and unconditional text generation. We make the following contributions:
In section 3, we introduce Sed, the first continuous diffusion approach for text with good scaling properties (testing models up to 420M parameters). We analyze several continuous text diffusion settings, and identify self-conditioning and diffusion on small fixed embeddings as key factors to make continuous text diffusion work.
In section 4, we apply classifier-free guidance (Ho & Salimans, 2022) to text data—an original achievement. We show that Sed can rival AR models on generic language tasks, for similar models sizes. Sed samples achieve a better likelihood-entropy trade-off compared to these models, and are deemed comparable (if slightly worse) by human raters.
Related work
We provide an overview of diffusion models with a focus on modeling discrete data, as well as AR models and sample-based metrics for evaluating text generation.
Continuous diffusion has recently established itself as the method of choice for modeling continuous data such as images. While our main focus in this paper is on discrete data, we review some key works in continuous data modeling as this literature was the major source of inspiration for Sed. The first continuous diffusion formulation was introduced in the seminal work by Sohl-Dickstein et al. (2015). Ho et al. (2020) improved and simplified this formulation, relating it to denoising score matching, and creating a new method called DDPM. Nichol & Dhariwal (2021) further improved upon DDPM, showcasing impressive diffusion results compared to GANs. Rombach et al. (2021, Stable Diffusion) introduced diffusion in latent space. Conceptually similar to Sed, it was specifically targeted at image modeling. Classifier-free guidance was proposed by Ho & Salimans (2022) as a mean to improve image fidelity at the cost of reduced diversity. GLIDE (Nichol et al., 2022) scaled up the ideas of guided diffusion, while DALL-E 2 (Ramesh et al., 2022) and Imagen (Saharia et al., 2022) are the latest, most advanced image generation systems to date, combining most of the improvements proposed in previous works.
Discrete diffusion on discrete data.
One cannot simply reuse the methods that are successful on continuous image data in the discrete text domain. A number of bespoke methods have been explored instead, forming the family of discrete diffusion approaches. In discrete diffusion, the data is corrupted by switching from one discrete value to another. This was first proposed in the seminal work by Sohl-Dickstein et al. (2015), where it was tested on simplistic binary heartbeat data. It was extended to multinomial text modeling (Hoogeboom et al., 2021) and further scaled up in the D3PM work (Austin et al., 2021). Most recently, a similar discrete diffusion approach was applied to image modeling in VQ-Diffusion (Gu et al., 2022). In parallel, a few diffusion-like approaches were proposed in the denoising autoencoders literature. CMLM (Ghazvininejad et al., 2019) tackled machine translation. SUNDAE (Savinov et al., 2022) was the first non-AR method to show strong results both in machine translation and unconditional text generation. MaskGIT (Chang et al., 2022) demonstrated excellent results in modeling VQ-discretized images. These approaches rely on training models to predict masked tokens from their context, and iterating this reconstruction step multiple times at sampling time. Despite those positive developments, the samples from discrete diffusion methods for text modeling remains less coherent than those produced by AR methods.
Continuous diffusion on discrete data.
Fewer works try to tackle diffusion on discrete data from the same angle as Sed – starting by turning the data into continuous representations before modeling it with continuous diffusion formulations. Mittal et al. (2021) used a VAE to generate such representations for discrete music modeling, with exciting results. Closest to Sed, Diffusion-LM (Li et al., 2022) trains a token embedding together with the diffusion model itself. Diffusion-LM meets success on specific language applications, in low data regime and on constrained, very formatted textual data. Most recently, Analog Bits (Chen et al., 2022) introduced self-conditioning, closely related to step-unrolls in SUNDAE (Savinov et al., 2022), together with bit-level modeling to improve the generation of discretized images. While the qualitative results of those continuous methods on text modeling show promise, they have not been shown to scale to large realistic text datasets like C4 (Raffel et al., 2020) yet, or to compare with AR approaches on generic language tasks.
Auto-regressive modelling on discrete data.
AR models remain the method of choice for modeling discrete data. In combination with neural networks, they were first explored by Bengio et al. (2000) and later combined with RNNs (Sutskever et al., 2011). Their breakthrough moment came with the advent of the Transformer architecture, introduced by Vaswani et al. (2017) for machine translation. Even more impressive results were shown with GPT-3 (Brown et al., 2020), which trained a large AR language model unconditionally, and used few-shot prompting to adapt it to new tasks. A few works later improved upon the results of GPT-3, including Hoffmann et al. (2022).
Sample-based evaluation of text generative models.
There are traditionally two classes of metrics for generative modeling: likelihood-based and sample-based. While the likelihood-based way is mathematically appealing, its usefulness for measuring progress is reduced by the fact that not all models readily provide likelihood computation. Just like the sampled-based FID metric was important for driving the progress of diffusion in image modeling, there is a need for a sample-based metric which would be universally accepted for text modeling. Caccia et al. (2018) investigated fidelity/variance metrics for evaluating text GANs. Semeniuta et al. (2018) suggested using FID for texts. De Masson d’Autume et al. (2019) later used those previously proposed metrics to iterate on ScratchGAN but did not provide conclusive guidance on which metric a practitioner should choose – essentially finding serious vulnerabilities in all investigated metrics. We opted for a middle ground, reporting both sample likelihood according to a strong AR model and human preferences.
Method
In this section, we outline the different components of Sed: continuous diffusion in the space of token embeddings and self-conditioning, which form the basis of our approach for unconditional text generation; span masking and guided diffusion to enable conditional generation.
This parametrization gives us a closed form to sample for any arbitrary , given :
where , , \epsilon_{t}\sim\mathcal{N}\big{(}0,{\bm{I}}\big{)} and \epsilon\sim\mathcal{N}\big{(}0,{\bm{I}}\big{)}.
We define our generative model by approximately inverting the diffusion process of Eq. 1 to obtain a reverse process. The reverse process starts from and is defined as a Markov chain with learned Gaussian transitions (parameterized by , the weights of a neural network): {\bm{x}}_{t-1}\sim p_{\theta}(\cdotp|{\bm{x}}_{t})=\mathcal{N}\big{(}{\bm{\mu}}_{\theta}({\bm{x}}_{t},t),\sigma(t)^{2}{\bm{I}}\big{)}. We train a neural network to predict an estimate of the data and approximate the reverse process by using the following parametrization, with learnable means but fixed variances, and a fixed schedule :
While there exists a tractable variational lower-bound (VLB) on , Ho et al. (2020) showed that better results are obtained by optimizing a simplified objective that re-weights the terms in the VLB. We follow this approach, which simplifies the loss to a sum of mean-squared errors between the ground truth data and its estimates :
Though this framework works out of the box on images, which are close to continuous, we cannot apply it directly to the discrete tokens of the text modality. To resolve this issue, we perform continuous diffusion in a continuous space in which we embed text tokens.
To train the readout step, we add a reconstruction loss to during training. Conveniently, it naturally arises when deriving the VLB of with this discretization step (Li et al., 2022), introducing a simple cross-entropy loss to maximise :
Equipped with these 3 components we can train models to generate text, though only unconditionally. To add conditional generation to our system’s capabilities, we use two additional methods.
2 Span masking and guidance for conditional text generation
By design diffusion models for text generation are flexible and can handle a wide variety of infilling tasks. This is a key advantage over the predominant auto-regressive language models that typically generate text in a left-to-right fashion.
Span masking. We train our model on a rich set of infilling tasks with the following method. We split between two set of tokens, diffusion tokens over which we apply diffusion and optimize the diffusion loss from Eq. 4, and conditioning tokens that remain fixed. Conditioning tokens are defined by a binary conditioning mask set to one on conditioning positions and zero on positions to be infilled.
This span masking strategy defines a collection of text generation tasks with a large variety of conditioning which on average evenly splits the sequence between conditioning and infilling spans. It enables conditional generation, and opens the door for additional diffusion improvements.
Experiments
We train all our models on the C4 dataset (Raffel et al., 2020), using a SentencePiece tokenizer (Kudo & Richardson, 2018) composed of 32000 words. We use a non-causal transformer model (Vaswani et al., 2017) as our diffusion model (see Appendix A for details). Sed models are trained with sequence length 256, while for ablations models are trained with sequence length 128. We insert uniformly, i.e. not necessarily at the end of the sequence, 10% of padding tokens in the training set to allow Sed models to generate samples of varying size and provide more flexibility.
To generate word embeddings, we train a BERT model of fixed size (m parameters) and feature dimension . The diffusion space is defined by the initial lookup table of this BERT model. We bottleneck the dimension of the word embeddings and add a linear projection layer from to at the beginning of the model. We found this helped diffusion (see section 4.4).
Sed models are trained with a cosine noise schedule (Dhariwal & Nichol, 2021), with , and . We use batches of 65.536 tokens, thus for sequence length 256 the batch size is set to 256. We use a maximum span count of 5 for all runs except for its specific ablation. We train Sed models at two different scales: Sed-S (m parameters, training steps) and Sed-L (m, steps). Their detailed architectures can be found in Appendix A.
2 Validation
While optimizing the perplexity of AR models for text leads to improved language models, directly optimizing the ELBO of diffusion models for images does not correlate strongly with sample quality as observed by Nichol & Dhariwal (2021); Kingma et al. (2021); Ho & Salimans (2022). For images, the sample based metric FID (Heusel et al., 2017) has been introduced as a measure of sample quality and is now widely adopted. Similarly, we need a sample-based metric for text generation that is reliable and allows comparison between a large variety of generative models. To provide a fair comparison to AR models, we rely on three metrics.
The first metric measures how likely the samples produced by a model are according to an AR language model with 70B parameters, trained on 1.4B tokens (Hoffmann et al., 2022); we denote this metric AR NLL for auto-regressive negative log-likelihood. It provides a continuous measure of sample quality that has proven useful when combined with a measure of sample diversity, e.g. in the development of nucleus sampling (Holtzman et al., 2020) for improved AR model decoding.
To measure diversity we rely on a second metric, the unigram entropy of samples, which helps balance the AR NLL that can be gamed by unnatural repetitive samples. For both these metrics, our target is the score of the validation set data. Deviating from the data unigram entropy in particular is a sign of degenerate modeling.
Though this initial combination has provided us with a reliable signal to iterate over our model design, it remains imperfect; it too can be gamed, though it is harder to do so. To address this limitation, we also report human preferences. We presented 6 colleagues with 20 pairs of samples for each comparison, asking them to pick the best one.
For all three metrics, we report results on two tasks: unconditional language modeling and suffix in-filling, the later a heavily conditioned task.
3 Results
Samples. We present samples generated with our Sed models in Table 10. We use a single model to perform a wide variety of text generation tasks, such as unconditional generation, filling-in-the-middle or filling several spans of text. We show strong performance in the unconditional case, with samples that are syntactically correct and stay coherent on long sequences. In the conditioned case, Sed models are able to infill spans with coherent transitions and links to the conditioning but also exhibit a rich diversity. By design, Sed yields flexible bi-directional masking models that can perform text generation on a diverse set of conditioned task. To compare Sed with AR baselines we next restrict conditioning to a prefix and consider a task of suffix in-filling.
Comparison to AR models. To assess the generation ability of Sed, we compare against AR baselines of similar capacity and trained following optimal scaling laws from Hoffmann et al. (2022) on suffix in-filling. We sample a batch of sequences from C4 and use the first 128 tokens as conditioning given to the model to generate a suffix of 128 tokens. Figure 1 reports AR NLL and unigram entropy of the generated suffixes for AR and Sed models. As a reference point, we compute the AR NLL and unigram entropy of the ground truth C4 suffixes and report it on the plot. Several methods can be used to improve sampling quality at the cost of samples diversity; we use nucleus sampling (Holtzman et al., 2020) for AR models and guidance (Dhariwal & Nichol, 2021; Ho & Salimans, 2022) for Sed models. We show the impact of guidance on samples quality in Table 3. To our knowledge, we are the first to show sample quality improvement when using guidance for text generation.
As shown in Figure 1, both Sed-S and Sed-L perform strongly when compared against AR baselines – even though we report a metric favoring AR models on a task AR models have been designed to optimize. Similar to nucleus sampling for AR models, guidance has a strong positive impact on sample quality that is both observed quantitatively with improved AR NLL in Figure 1 and qualitatively in Table 3. We observe that using a top- nucleus sampling below for AR models or a guidance scale above for Sed models leads to samples exhibiting a lot of repetitions, a degenerate case reflected by a lower entropy of samples even though sample AR NLL improves.
Our human preference scores temper our observations in Table 4. They show that our NLL and entropy metrics do not tell the whole story, as humans still prefer AR models at equivalent size. While Sed-L performs slightly worse than AR-L (38% preference in suffix in-filling, 44% on unconditional generation), its scores remain comparable. Sed-L is roughly on par with AR-S.
Finally, we compare Sed and AR models’ qualitative examples with short prompts in Table 2.
4 Ablations
Self-conditioning and embedding pretraining. Results from Table 5 and samples from Table 6 show the influence of both the diffusion space and self-conditioning. AR NLL decreases very significantly when using self-conditioning, regardless of the rest of the setup. Diffusing at the bit-level (Chen et al., 2022) yields very high NLLs. While using random embeddings performs markedly better, using pretrained embeddings results in further improved numbers.
Samples from Table 6 highlight that models trained on random word embeddings exhibit topic modelling abilities with the co-occurrence of words like child and mother even though the paragraph remains globally incoherent and meaningless tokens like gluc are generated. Self-conditioning dramatically improves sample quality; the diffusion model gets the low-level structure right and generates syntactically correct sentences, even though the global text is not intelligible. Combining self-conditioning and pretrained embeddings leads to globally coherent paragraphs that stay on topic with proper sentence structure.
Embedding dimension. An important design choice for SED is the word embeddings space. We study the influence of pretrained embedding size in Table 8. Surprisingly, there is a threshold after which performance degrades when increasing the dimension of embeddings. We visualize the forward process for different embedding sizes by displaying the nearest neighbor of a noised token while running the forward process. In high dimension we observe that the nearest neighbor of a noised token remains the starting token itself until it switches to a completely random, unrelated token. In low dimension, we often observe that the closest neighbor of a noised token goes through several semantically related tokens (nearest neighbor of the starting token) before ultimately becoming random. We hypothesize that the random walks defined by diffusion are more likely to drift towards neighbors of the starting token in low dimension. As a result, when diffusing in low dimension information is destroyed in a more semantically meaningful fashion, which leads to an easier learning problem for the denoising function.
Number of spans. In order to enable in-filling, we train the model not only to do unconditional generation but also to conditionally fill spans of tokens. For each data point we sample a span number uniformly at random and span delimiters to generate the span mask. Picking the maximum allowable number of spans has a significant effect on model performance, as we can see in Table 8. Somewhat counter-intuitively, adding span masking improves even unconditional generation NLLs. It also appears that using a relatively high maximum span number is optimal. We hypothesize that this results in a varied mix of task difficulty at training time, between ”easy”, very conditioned problems on the one hand and ”harder”, unconditional ones on the other.
Scaling. We show encouraging results when scaling from Sed-S (150m) to Sed-L (m). We train both models on sequences of tokens and report a AR NLL of for Sed-S compared to for Sed-L. This improvement translates to improved sample quality, as is confirmed by our human preference scores, which are much higher for the larger model (63%, see Table 4).
Limitations
While our results are promising and show that continuous diffusion for text can be an exciting alternative to AR models, the current approach does present some significant limitations.
First, much more could be done in terms of model tuning, including scaling to much bigger models to better understand Sed’s limits, and to be able to compare it with state-of-the-art AR models. Our training regime in particular would certainly benefit from more hyperparameter optimisation.
Second, one compelling reason we chose to explore continuous diffusion for text is to leverage the improvements produced by the literature on image generation. While we have ported some (e.g. self-conditioning), a lot more remains unexplored. The most obvious example is the sampling process itself, where the number of required steps has been considerably reduced for images (e.g. Karras et al. (2022) goes from 1000 to 35, and Salimans & Ho (2022) all the way down to 4 on simple images). Our current sampling is very inefficient, and this direction is one of the first improvements to make over Sed.
Third, Sed crucially relies on diffusing in a pretrained embedding space. This means relying on a second model, and using embeddings that may not be optimal for diffusion. Ideally, we’d train the full model end-to-end, which could yield even better results. While Li et al. (2022) found some success with this approach, it was in a specific setting at a small scale; in practice we found it difficult to avoid competition between the diffusion and reconstruction loss.
Finally, our work would benefit from improved metrics in the experimental section. Because the current state of the art involves AR models, the field lacks established benchmarks for tasks diffusion models are potentially better suited for, such as text in-filling. We opted for a reasonable mix, evaluating the negative log-likelihood of generated samples according to a very strong AR model as well as their token entropy and complementing it with a human evaluation. However, both NLL and unigram entropy are gameable (e.g. AR models assign very low NLL to repetitive snippets, and long enough repetitions can fool even entropy). Further, our NLL is inherently tied to its AR model and could thus be providing an unfair advantage to AR models. All told, we still found both metrics quite useful for measuring research progress, and our human evaluation confirmed our results. Moving forward, defining a clean in-filling benchmark would help produce even more convincing results.
Conclusion
We propose Sed, the first generally-capable continuous diffusion model for text generation. Sed models can perform both conditional and unconditional generation, and their performance rivals AR models while being more flexible in their use (e.g. enabling in-filling). We demonstrate their performance and study the impact of the main design choices.
Despite its limitations, this work lays the foundation for more exciting research. Promising directions include speeding up the sampling following the lessons learnt in the image domain, devising better embedding spaces for diffusion and investigating new in-filling capabilities.
References
Appendix A Model architecture
Appendix B Forward diffusion process visualization
To support the discussion on word embeddings dimension from Section 4.4, we present a visualization of the forward diffusion process. Given starting tokens , we project the noised tokens of the forward process at step to their nearest neighbor among word embeddings to obtain . We then store the 128 nearest neighbors of starting tokens and define the rank of at its index in . We display and highlight it in green if is close to zero (meaning is a close neighbor of ) and in increasingly red colors otherwise. We present the first 16 nearest neighbors of in Figure 2 and provide an illustration of the color code used for highlighting. Figure 3 shows an instance of the forward diffusion process while diffusing on embeddings of with a high dimension of 896 and Figure 4 shows diffusion on embeddings with a lower dimension of 32.
In contrast, in lower dimension we see meaningfully-related tokens appear as the corruption progresses (‘brown’ becomes ‘grey’, ‘quick’ becomes ‘swift’, ‘over’ becomes ‘underneath’ etc). We believe this more gradual information destruction is beneficial for the diffusion model.