AR-Diffusion: Auto-Regressive Diffusion Model for Text Generation
Tong Wu, Zhihao Fan, Xiao Liu, Yeyun Gong, Yelong Shen, Jian Jiao, Hai-Tao Zheng, Juntao Li, Zhongyu Wei, Jian Guo, Nan Duan, Weizhu Chen
Introduction
Text generation is a fundamental task within the field of natural language processing (NLP). Pre-trained language models like GPT-4 (OpenAI, 2023), LLaMA (Touvron et al., 2023), and Alpaca (Taori et al., 2023) have garnered significant attention with their ability to generate fluent and human-like textual content. These models utilize the auto-regressive (AR) Transformer decoders (Vaswani et al., 2017) to emit generated tokens one-by-one in sequential order from left to right. By leveraging the power of position dependency, AR models are able to enhance the naturalness, coherence, and adherence to human language conventions in the generated text (Brown et al., 2020).
Recent studies have shown the remarkable performance of diffusion models in image generation (Ho et al., 2020), motivating researchers to extend diffusion to text generation (Li et al., 2022a; Gong et al., 2022; Dieleman et al., 2022; Yuan et al., 2022; Ye et al., 2023). By introducing timestep, these methods progressively regulate the interpolation between the original tokens and Gaussian noise, then iteratively denoise for text generation. At each timestep, the diffusion-based text generator predicts all tokens simultaneously following Non-Auto-Regression (NAR) (Lewis et al., 2020; Qi et al., 2020, 2021; Li et al., 2022b), leading to faster decoding speed compared to AR. However, it also inherits the drawback of NAR, namely the sacrifice of inter-token position dependency (Li et al., 2022c) and the drop of generation performance (Bao et al., 2021).
To conduct a comprehensive analysis, we introduce a two-dimensional coordinate system to track the diffusion timestep of tokens positioned at various locations. As illustrated in Figure 1, the system assigns the token position to the horizontal axis and the diffusion timestep to the vertical axis. Diffusion-LM (Li et al., 2022a), which is followed by existing diffusion-based text generation models, is shown in Figure 1(a). It assigns a uniform timestep to all tokens. In contrast, tokens in the AR model depicted in Figure 1(b) exhibit distinct timesteps within a generation step (). For instance, the already decoded token at position has a timestep of , while the to-be-decoded token at position has a timestep of . This approach effectively captures the sequential dependency. Motivated by this observation, we introduce AR-Diffusion, an auto-regressive diffusion method, for the disparity in token positions and the principle of sequential token identification.
In AR-Diffusion, we propose a multi-level diffusion strategy that includes both sentence-level and token-level diffusion. We randomly choose a sentence-level timestep , and assign dynamic movement speeds by determining position-sensitive token-level timestep for each token. This enables tokens at the left of a sentence to undergo faster movement from random Gaussian noise to token embedding, while those at the right of the sentence experience slower movement to better utilize information from previously denoised tokens. During inference, to reduce the significant number of inference steps (e.g., 2,000) required in Diffusion-LM (Li et al., 2022a), SeqDiffSeq (Yuan et al., 2022) and GENIE (Lin et al., 2023), we introduce a skipping mechanism that collaborates with the multi-level diffusion strategy to accelerate the process.
Experimental results across various text generation tasks, such as text summarization, machine translation, and common sense generation, have consistently demonstrated that AR-Diffusion surpasses existing text diffusion models, including AR methods in terms of both quality and diversity. Moreover, our verification reveals that AR-Diffusion requires fewer resources during decoding while maintaining superior performance. It achieves faster than SeqDiffSeq (Yuan et al., 2022) in machine translation and faster than GENIE (Lin et al., 2023) in text summarization while delivering comparable results. Furthermore, it demonstrates promising results even in a challenging scenario where decoding is limited to only two steps.
Preliminary
In the field of natural language generation, conditional generative models are commonly implemented using either auto-regressive (AR) or non-auto-regressive (NAR) methods. In AR (Vaswani et al., 2017), tokens on the right are predicted based on visible left tokens. The likelihood is given by , where denotes the -th token of . On the other hand, NAR (Gu et al., 2017) assumes conditional independence among tokens and generates them uniformly without distinction during decoding, resulting in the likelihood . This parallel generation approach is of lower quality compared to AR, although it offers a substantial speed advantage.
2 Diffusion Models for Text Generation
Recently, Li et al. (2022a) propose a natural language generation model based on the diffusion process, which is typically divided into a forward noising process and a reverse denoising process.
Specifically, the forward process is a fixed linear Gaussian model, which gradually perturbs the random variable until it becomes the standard Gaussian distribution. This can be formalized as:
where, , and is a coefficient that monotonically decreases with timestep , is the latent state at timestep .
The reverse process is to initiate from standard Gaussian noise and progressively utilize the denoising transition for generation.
where the mean and variance are learned from the model. In particular, we follow Li et al. (2022a)’s approach of using predefined variance without trainable parameters.
In consequence, combined with maximizing the evidence lower bound (ELBO) of , our training objective of the conditional diffusion language model is:
Methodology
In the typical diffusion process, every token in the text sequence has the same diffusion timestep. In order to leverage the sequential nature of language, we enable tokens to have different diffusion timesteps during the forward and reverse pass. To accomplish this, we propose a multi-level diffusion strategy that includes both sentence-level and token-level diffusion.
Firstly, at the sentence-level, we follow Diffusion-LM (Li et al., 2022a) to randomly select a timestep . Secondly, at the token-level, we incorporate positional information based on the sentence-level timestep to regulate the diffusion timestep for the current token. The procedure is illustrated as:
where is the given target sentence length, is the sentence representation at timestepPlease note that if we talk about a “timestep” without explicitly indicating that it is for token-level, it should be for sentence-level. , is the latent representation for the -th token at sentence-level timestep , and is a token-level timestep function that denotes the token-level diffusion timestep determined by token position and sentence-level timestep .
We visualize the token-level timestep \big{(}n,f(n,t)\big{)} onto a two-dimensional coordinate system as Figure 1 , which takes the token position as the horizontal axis and the sentence-level timestep as the vertical axis. Furthermore, to provide a more profound description of the characteristics of movement, we define the speed of movement as the following equation.
where and are the start and end sentence-level diffusion timesteps. It can be observed that tokens in Diffusion-LM share the same movement speed, while those in AR exhibit different speeds.
2 Token-Level Diffusion with Dynamic Movement Speed
Based on the speed of movement, we propose a fundamental principle, dynamic movement speed, for designing the token-level diffusion timestep function to take advantage of AR in diffusion. Specifically, elements on the left side of a sentence undergo higher movement speed from random Gaussian noise to token embedding, while those on the right side experience lower movement speed, thereby they can be generated in the later sentence-level timestep and utilize information from previously generated tokens more effectively.
In the reverse diffusion process, the multi-level diffusion follows the formula:
where denotes the -th element.
3 Inference with Skipping
Typically, the generation process needs to go through all the sentence-level timesteps from to . To reduce the decoding time, we introduce a skipping mechanism that allows us to traverse a subset of timesteps.
To ensure consistency between training and inference, we also need to calculate the timestep for each token during the inference process. Therefore, we first establish an anchor point, and then uniformly select a decreasing subsequence from all timesteps ( to ). The count of this sequence is the total decoding steps (). For example, assuming the interval is and is , then is , and the subsequence is $$.
Each element of this subsequence represents the sentence-level timesteps , and we can use equation 6 to calculate . Then, based on equation 7, we calculate the token-level timesteps corresponding to each position. We take the current sentence-level timestep and the next sentence-level timestep , and calculate the token-level timesteps and for each position. Since , , implying that . The essence of Skipping is reflected in the fact that each token undergoes significant span during denoising (from to ).
In practice, we propose an algorithm for the inference, illustrated in Algorithm 2.
In equation 10, the conditional distribution of is inferred by , and then we decompose it by positions due to the independent forward process of elements at different positions. From equation 11 to equation 12, we establish the relationship between tokens at different timesteps, and the detailed derivation can be found in Appendix A.
Experiments
This task involves taking a long document as input and generating a concise sentence as output. This requires models with the ability to identify important content and rewrite it in a condensed form. In our experiments, we use the publicly available XSum (Narayan et al., 2018) and Cnn/DailyMail Hermann et al. (2015) on GLGEhttps://microsoft.github.io/glge/, which is also named as GLGE-Easy.
Machine Translation
Translation is a widely used sequence-to-sequence task. The input is a sequence of words in the source language, and the output is a sequence of corresponding words in the target language. We choose the IWSLT 2014 dataset and the data processing method is to follow the scripts provided by fairseqhttps://github.com/facebookresearch/fairseq/tree/main/examples/translation.
Common Sense Generation
In this task, the model is provided with a concept set consisting of objects and actions as input. The objective is to generate a sentence that incorporates these concepts and describes a realistic scenario. We use CommonGenhttps://inklab.usc.edu/CommonGen/ dataset for evaluation.
2 Experimental Details
Our model configuration is implemented based on Transformer-base (Vaswani et al., 2017). In particular, For XSum and Cnn/DailyMail, we set the diffusion embedding dimension to 128. For IWSLT14, we use 64-dimensional diffusion embedding, 4 attention heads and 1024-dimensional feed-forward layers. For CommonGen, we adopt 64-dimensional diffusion embedding, 8 attention heads and 512-dimensional feed-forward layers.
Training and Inference
In the training phase, we employ a square-root noise schedule and 2,000 diffusion steps (Li et al., 2022a). Specially, we use the tokenizer and vocabulary constructed by Byte Pair Encoding (BPE)We train bpe on the training set, and follow the vocabulary size of fairseq, IWSLT14 is set to 10,000 . (Kudo and Richardson, 2018) for translation tasks. For other tasks, we adopt the tokenizer and vocabulary of bert-base-uncased.
Baselines
NAR: NAT (Gu et al., 2017), iNAT (Lee et al., 2018), CMLM (Ghazvininejad et al., 2019), LevT (Gu et al., 2019) and CNAT (Bao et al., 2021);
Semi-NAR: InsT (Stern et al., 2019), iNAT (Lee et al., 2018), CMLM (Ghazvininejad et al., 2019) and LevT (Gu et al., 2019);
AR: bRNN (Gu et al., 2016), LSTM (Greff et al., 2017) and Transformer (Vaswani et al., 2017);
Diffusion: DiffusionLM (Li et al., 2022a), CDCD (Dieleman et al., 2022), SeqDiffuSeq (Yuan et al., 2022), DINOISER (Ye et al., 2023) and GENIE (Lin et al., 2023).
Metrics
We follow the approach of Qi et al. (2020)https://github.com/microsoft/ProphetNet/tree/master/GLGE_baselines to evaluate the ROUGE-1/2/L of the summarization task. For the evaluation of translation tasks, we adopt the setting of SeqDiffuSeq (Yuan et al., 2022) to report BLEU score. In addition, we also calculate the SacreBLEU score according to the setting of DINOISER (Ye et al., 2023) for comparison. For CommonGen, we employ ROUGE-2/L, BLEU-3/4, METEOR and SPICE under the evaluation methods of Lin et al. (2020)https://github.com/INK-USC/CommonGen/tree/master/evaluation/Traditional/eval_metrics.
Training Parameters
Our training parameters on different datasets are shown in Table 1. Our linear schedule warm up steps is 4,000 , where denotes gradient accumulation number. In addition, we use the AdamW (weight decay = 0.0) optimizer and dropout is 0.2. All experiments are implemented on 8 Tesla V100-32G. It takes about 20 hours to train XSum and Cnn/DailyMail, about 5 hours to train IWSLT14, and about 2 hours to train CommenGen.
3 Main Results
The results on different datasets are shown in Table 2, Table 3, Table 4 and Table 6. The best result is bolded and the second-best result is underlined . “” indicates the number of generated candidate samples. It can be seen from the results in each table that AR-Diffusion achieves the best performance.
During the inference process, we utilize 20 inference steps and employ Minimum Bayes Risk (MBR) (Kumar and Byrne, 2004) decoding to select the best sample, following (Li et al., 2022a). We choose MBR instead of the selection approach in GENIE, as GENIE picks up the best sample by calculating the maximum score for each generated one using ground truth, which introduces unfairness. To ensure a fair comparison, we re-implement GENIE using our configuration and perform inference with 20 steps.
The results presented in Table 2 and Table 3 clearly demonstrate that AR-Diffusion outperforms the existing NAR and Semi-NAR approaches across all metrics. Moreover, AR-Diffusion consistently achieves significant improvements over GENIE in terms of all indicators. Furthermore, in comparison to Transformer, AR-Diffusion outperforms it on both ROUGE-1 and ROUGE-L, while achieving comparable performance in terms of ROUGE-2. Notably, when the sample number is 500, AR-Diffusion demonstrates superiority over Transformer across all the measures.
Machine Translation
Table 4 presents the BLEU score implemented by SeqDiffuSeq setting. AR-Diffusion outperforms the non-auto-regressive CNAT in greedy search for a single sample, and achieves a substantial gain. Moreover, the BLEU score of AR-Diffusion surpasses GENIE by a large margin and shows a slightly better performance than the AR Transformer. Specially, AR-Diffusion achieves a more powerful result at = 500.
In Table 5 we give the SacreBLEU score according to the setting of DINOISER. AR-Diffusion has notable improvements over non-auto-regressive CMLM. Besides, AR-Diffusion achieves excellent performance among text diffusion models for both En→De and De→En tasks. Specifically, AR-Diffusion is far superior to GENIE and comparable to the newly proposed DINOISER at = 50. Nevertheless, the performance is stronger than DINOISER when = 500DINOISER has shown in their Figure 4 that their method is not better with a larger ..
Common Sense Generation
As depicted in Table 6, AR-Diffusion achieves superior performance compared to the current AR, NAR, and other diffusion methods across all the metrics on the CommonGen dataset.
4 Inference Efficiency
First, we use the number of function evaluations (NFE) as a measure to compare inference efficiency (Ye et al., 2023) in machine translation. From Table 4, it is evident that even when the NFE is reduced to 1% of SeqDiffuSeq (equivalent to faster), AR-Diffusion still outperforms SeqDiffuSeq. Moreover, increasing the number of generated candidate samples () leads to further performance improvements, albeit with increased time consumption.
Second, we conduct experiments with an extremely limited number of inference steps (2 and 3)The time consumed by each step in the inference process is exactly the same. and compare the performance with that of GENIE in XSum. The results are presented in Table 7. When reducing the number of steps to 2, GENIE experiences a significant decline, with an average score of 4.20 in the AVG Drop column, while AR-Diffusion exhibits a comparatively smaller decrease of 1.34. Furthermore, with 3 steps, although the performance deterioration of GENIE is somewhat reduced, the average score still shows a decline of 2.81. In contrast, AR-Diffusion maintains a high performance level, with an average score differing from the 20-step result by only 0.64. Notably, the results of AR-Diffusion at 3 steps are comparable to the results of GENIE at 2,000 steps. Therefore, compared to GENIE, the inference speed of AR-Diffusion can be accelerated by up to .
5 Analysis
Diversity is a key advantage of diffusion models. To measure the diversity of generated samples, We adopt the SELF-BLEU (Zhu et al., 2018) metric, in which a lower score indicates higher diversity. In Lin et al. (2023), various sampling methods were applied to the pre-trained auto-regressive model BART, including Greedy Search, Beam Search Xiao et al. (2022), Diverse Beam Search(diversity strength = 0.8) Vijayakumar et al. (2016), Typical Sample ( = 1.2) Meister et al. (2022), Top-k Sample ( = 50) Fan et al. (2018) and Nucleus Sample ( = 0.92) Holtzman et al. (2020).
Specifically, greedy search is to select the token with the highest probability at each step. Beam search is to select the largest token from among the beams with higher probability at each step. Diverse beam search is to divide the beams into multiple groups at each step and ensure the difference between groups by calculating the diversity score between groups. Typical sampling selects samples through a discrete random process. Top-k sampling is to randomly select one of the candidate tokens with the highest probability at each step. Nucleus sampling is to randomly select one token at each step from the candidate tokens whose probability density is greater than .
As shown in § 4.4, AR-Diffusion achieves significantly higher diversity compared to the auto-regressive model. Furthermore, the diversity can be comparable to GENIE with a better performance.
Ablation Study
To demonstrate the effectiveness of our proposed method, we perform ablation experiments on the XSum dataset. Our results show that both our proposed multi-level diffusion and skipping mechanism are essential for achieving the high performance of AR-Diffusion.
Maintaining the skipping inference method, we remove the token-level diffusion during the training process, which degenerates to GENIE w/ skipping. The comparison results are shown in Figure 2. It can be observed that after removing, the AVG-ROUGE score is greatly lower after 2 steps.
The performance of applying our proposed skipping mechanism and DDIM (Song et al., 2021) to AR-Diffusion is shown in Figure 2. The results demonstrate that the skipping mechanism consistently outperforms DDIM in various inference steps. Additionally, the skipping mechanism can be easily applied to GENIE. As depicted in Figure 2, DDIM suffers a significant drop in performance when the number of inference steps is less than 40. In contrast, the skipping mechanism consistently maintains good performance across all inference steps.
Case Study
By mapping the state to the token with the highest logits, we visualize the intermediate states of AR-Diffusion. As depicted in Figure 3, AR-Diffusion undergoes a denoising process, transforming the random Gaussian noise into a coherent sentence over 20 steps, and we present 5 of them. With the progression of each timestep, compared to the tokens on the right side of the sentence, the tokens on the left side demonstrate faster determination and a rapid increase in the corresponding logits. This behavior is consistent with our principle of dynamic movement speed from left to right.
6 Impact of Minimum Bayes Risk and Anchor Point
To investigate the relationship between the number of generated candidate samples () and the quality of generation, we generate varying numbers of samples, ranging up to 1,000, on the IWSLT14 De→En test set and present the results in Figure 4. The curve demonstrates an initial gain of approximately 0.5 SacreBLEU within the first 200 samples, after which the gain becomes insignificant with generating more samples.
Anchor Point
We conduct experiments on AR-Diffusion using different anchor points . These anchor points vary in terms of values, namely , and , where denotes the target sentence length. Additionally, they share a common value of , which represents the total time step of diffusion. We present the results in Figure 4, and determine that the best result is achieved at = .
Related Work
AR models have been the dominant approach for text generation (OpenAI, 2023; Touvron et al., 2023; Dong et al., 2023), but their token-by-token generation nature often leads to unsatisfactory inference speed. To address this issue, NAR models have been developed in recent years. The NAR method is initially proposed by Gu et al. (2017), its objective is generate the entire output sequence in parallel, thereby improving generation speed and efficiency. Subsequently, LevT (Gu et al., 2019) adopts insertion and deletion to address the lack of flexibility in NAR generation, CMLM (Ghazvininejad et al., 2019) utilizes a masked language model to improve the quality of NAR generation through a constant number of iterations, and CNAT (Bao et al., 2021) introduces latent variables to represent the category information of the target word to make full use of the latent representation. However, these NAR methods are hard to model inter-token position dependency and deficient in generation performance.
Continuous Text Diffusion
The application of diffusion models to continuous text space is first introduced by Li et al. (2022a). Through the embedding and rounding processes, the direct integration of continuous noise into word embeddings was accomplished. After that, more people attempt to adopt continuous text diffusion model to solve sequence-to-sequence tasks. DiffuSeq (Gong et al., 2022) divides the input into two parts, utilizing one part as a condition, and perturbs the other part with noise. CDCD (Dieleman et al., 2022) proposes score interpolation and time warping to allow diffusion model and Euclidean embedding to share the same loss function for training. SeqDiffuSeq (Yuan et al., 2022), GENIE (Lin et al., 2023) and DINOISER (Ye et al., 2023) incorporate diffusion model into the encoder-decoder structure through cross-attention mechanisms.
It is important to highlight the differences between our method and both ARDMs (Hoogeboom et al., 2022) and TimeGrad (Rasul et al., 2021), despite the common references to autoregression and diffusion in all these. ARDMs employ an order-agnostic technique, leveraging masking and prediction for generation in arbitrary orders. On the other hand, TimeGrad integrates RNN and diffusion to model the conditional distribution of future steps of multivariate time series. In contrast, our research focuses on implementing the diffusion process within a continuous embedding space, with the primary aim of generating text in a left-to-right sequence.
Conclusion
This paper introduces AR-Diffusion, which exhibits AR-like generation behavior but enables efficient parallel decoding. Embracing the inherent sequential nature of language, we propose a multi-level diffusion model, consisting of sentence-level and token-level components, to assign dynamic movement speeds to tokens. Consequently, compared to those on the right, the left tokens undergo fewer denoising steps and generate earlier to subsequently influence the later ones. Furthermore, we introduce a skipping mechanism to facilitate parallel generation within the multi-level diffusion framework. The experimental results across various tasks demonstrate that AR-Diffusion surpasses existing diffusion models in terms of quality while maintaining diversity. Additionally, compared to existing diffusion language models, AR-Diffusion achieves comparable results while being faster.
Limitation
A primary limitation of our work lies in the requirement of generating a large number of candidate samples for optimal performance. As an illustration in Table 3 of Cnn/DailyMail dataset, AR-Diffusion ( = 50) achieves a 0.8 lower ROUGE-2 score compared to AR-Diffusion ( = 500). We anticipate exploring more efficient sampling strategies to minimize the number of generated samples without performance drop.
References
Appendix A Proof of Inference with Skipping
During the inference process, skipping strategy requires the model to infer the state at a far-off timestep compared to the current state , where . In our model, due to the dynamic speed setting, token with smaller timestep , which is closer to , and positions can provide stronger auxiliary information than . This reduces the difficulty of inferring states for tokens in the end, making our multi-level diffusion model particularly suitable for accelerating the generation process.
Through maximizing the evidence lower bound (ELBO) of , the training object is equivalent to minimize the divergence between and following [Luo, 2022].
By converting the joint probability distribution into a conditional probability distribution, we obtain the following formula for .
Similarly, we reach the same conclusion regarding .
Based on equation 13, which consists of , and the interchangeability between and , we can decompose by incorporating and , and utilize our estimated to determine the expression of .
Next, we obtain the explicit expression q\big{(}z^{n}_{f(n,t_{i+1})}\mid z^{n}_{f(n,t_{i})},z_{0}^{n}\big{)} through linear interpolation between and .
where we have the following notations for simplification.
Building upon equation 15, we substitute with , yielding the final formula for p_{\theta}\big{(}{\bm{z}}^{n}_{f(n,t_{i+1})}\mid{\bm{z}}^{n}_{f(n,t_{i})};{\bm{x}}\big{)} as the following equation.