One Reference Is Not Enough: Diverse Distillation with Reference Selection for Non-Autoregressive Translation
Chenze Shao, Xuanfu Wu, Yang Feng
Introduction
Non-autoregressive machine translation Gu et al. (2018) has received increasing attention in the field of neural machine translation for the property of parallel decoding. Despite the significant speedup, NAT suffers from the performance degradation compared to autoregressive models Bahdanau et al. (2015); Vaswani et al. (2017) due to the multi-modality problem: the source sentence may have multiple correct translations, but the loss is calculated only according to the reference sentence. The multi-modality problem will cause the inaccuracy of the loss function since NAT has no prior knowledge about the reference sentence during the generation, where the teacher forcing algorithm Williams and Zipser (1989) makes autoregressive models less affected by feeding the golden context.
How to overcome the multi-modality problem has been a central focus in recent efforts for improving NAT models Shao et al. (2019, 2020, 2021); Ran et al. (2020); Sun and Yang (2020); Ghazvininejad et al. (2020); Du et al. (2021). A standard approach is to use sequence-level knowledge distillation Kim and Rush (2016), which attacks the multi-modality problem by replacing the target-side of the training set with the output from an autoregressive model. The distilled dataset is less complex and more deterministic Zhou et al. (2020), which becomes a default configuration of NAT. However, the multi-modality problem in the distilled dataset is still nonnegligible Zhou et al. (2020). Furthermore, the distillation requires NAT models to imitate the behavior of a specific autoregressive teacher, which limits the upper bound of the model capability and restricts the potential of developing stronger NAT models.
In this paper, we argue that one reference is not enough and propose diverse distillation with reference selection (DDRS) for NAT. Diverse distillation generates a dataset containing multiple reference translations for each source sentence, and reference selection finds the reference translation that best fits the model output for the training. As illustrated in Figure 1, diverse distillation provides candidate references “I must leave tomorrow" and “Tomorrow I must leave", and reference selection selects the former which fits better with the model output. More importantly, NAT with DDRS does not imitate the behavior of a specific teacher but learns selectively from multiple references, which improves the upper bound of the model capability and allows for developing stronger NAT models.
The object of diverse distillation is similar to the task of diverse machine translation, which aims to generate diverse translations with high translation quality Li et al. (2016); Vijayakumar et al. (2018); Shen et al. (2019); Wu et al. (2020); Li et al. (2021). We propose a simple yet effective method called SeedDiv, which directly uses the randomness in model training controlled by random seeds to produce diverse reference translations without losing translation quality. For reference selection, we compare the model output with all references and select the one that best fits the model output, which can be efficiently conducted without extra neural computations. The model learns from all references indiscriminately in the beginning, and gradually focuses more on the selected reference that provides accurate training signals for the model. We also extend the reference selection approach to reinforcement learning, where we encourage the model to move towards the selected reference that gives the maximum reward to the model output.
We conduct experiments on widely-used machine translation benchmarks to demonstrate the effectiveness of our method. On the competitive task WMT14 En-De, DDRS achieves 27.60 BLEU with speedup and 28.33 BLEU with speedup, outperforming the autoregressive Transformer while maintaining considerable speedup. When using the larger version of Transformer, DDRS even achieves 29.82 BLEU with only one decoding pass, improving the state-of-the-art performance level for NAT by over 1 BLEU.
Background
Gu et al. (2018) proposes non-autoregressive machine translation to reduce the translation latency through parallel decoding. The vanilla-NAT models the translation probability from the source sentence to the target sentence as:
where is a set of model parameters and is the translation probability of word in position . The vanilla-NAT is trained to minimize the cross-entropy loss:
The vanilla-NAT has to know the target length before constructing the decoder inputs. The target length is set as the reference length during the training and obtained from a length predictor during the inference. The target length cannot be changed dynamically during the inference, so it often requires generating multiple candidates with different lengths and re-scoring them to produce the final translation Gu et al. (2018).
The length issue can be overcome by connectionist temporal classification (CTC, Graves et al., 2006). CTC-based models usually generate a long alignment containing repetitions and blank tokens. The alignment will be post-processed by a collapsing function to recover a normal sentence, which first collapses consecutive repeated tokens and then removes all blank tokens. CTC is capable of efficiently finding all alignments which the reference sentence can be recovered from, and marginalizing the log-likelihood with dynamic programming:
Due to the superior performance and the flexibility of generating predictions with variable length, CTC is receiving increasing attention in non-autoregressive translation Libovický and Helcl (2018); Kasner et al. (2020); Saharia et al. (2020); Gu and Kong (2020); Zheng et al. (2021).
2 Sequence-Level Knowledge Distillation
Sequence-level Knowledge Distillation (SeqKD, Kim and Rush, 2016) is a widely used knowledge distillation method in NMT, which trains the student model to mimic teacher’s actions at sequence-level. Given the student prediction and the teacher prediction , the distillation loss is:
where are parameters of the student model and is the output from running beam search with the teacher model. The teacher output is used to approximate the teacher distribution otherwise the distillation loss will be intractable.
The procedure of sequence-level knowledge distillation is: (1) train a teacher model, (2) run beam search over the training set with this model, (3) train the student model with cross-entropy on the source sentence and teacher translation pairs. The distilled dataset is less complex and more deterministic Zhou et al. (2020), which helps to alleviate the multi-modality problem and becomes a default configuration in NAT models.
3 Diverse Machine Translation
The task of diverse machine translation requires to generate diverse translations and meanwhile maintain high translation quality. Assume the reference sentence is and we have multiple translations , the translation quality is measured by the average reference BLEU (rfb):
and the translation diversity is measured by the average pairwise BLEU (pwb):
Higher reference BLEU indicates better translation quality and lower pairwise BLEU indicates better translation diversity. Generally speaking, there is a trade-off between quality and diversity. In existing methods, translation diversity has to be achieved at the cost of losing translation quality.
Approach
In this section, we first introduce the diverse distillation technique we use to generate multiple reference translations for each source sentence, and then apply reference selection to select the reference that best fits the model output for the training.
The objective of diverse distillation is to obtain a dataset containing multiple high-quality references for each source sentence, which is similar to the task of diverse machine translation that aims to generate diverse translations with high translation quality. However, the translation diversity is achieved at a certain cost of translation quality in previous work, which is not desired in diverse distillation.
Using the randomness in model training, we propose a simple yet effective method called SeedDiv to achieve translation diversity without losing translation quality. Specifically, given the desired number of translations , we directly set different random seeds to train translation models, where random seeds control the random factors during the model training such as parameter initialization, batch order, and dropout. During the decoding, each model translates the source sentence with beam search, which gives different translations in total. Notably, SeedDiv does not sacrifice the translation quality to achieve diversity since random seeds do not affect the expected model performance.
We conduct the experiment on WMT14 En-De to evaluate the performance of SeedDiv. We use the base setting of Transformer and train the model for 150K steps. The detailed configuration is described in section 4.1. We also re-implement several existing methods with the same setting for comparison, including Beam Search, Diverse Beam Search Vijayakumar et al. (2018), HardMoE Shen et al. (2019), Head Sampling Sun et al. (2020) and Concrete Dropout Wu et al. (2020). We set the number of translations , and set the number of heads to be sampled as 3 for head sampling. We also implement a weaker version of our method SeedDiv-ES, which early stops the training process with only of total training steps. We report the results of these methods in Figure 2.
It is surprising to see that SeedDiv achieves outstanding translation diversity besides the superior translation quality, which outperforms most methods on both translation quality and diversity. Only HardMoe has a better pairwise BLEU than SeedDiv, but its reference BLEU is much lower. The only concern is that SeedDiv requires a larger training cost to train multiple models, so we also use a weaker version SeedDiv-ES for comparison. Though the model performance is degraded due to the early stop, SeedDiv-ES still achieves a good trade-off between the translation quality and diversity, demonstrating the advantage of using the training randomness controlled random seeds to generate diverse translations. Therefore, we use SeedDiv as the technique for diverse distillation.
2 Reference Selection
After diverse distillation, we obtain a dataset containing reference sentences for each source sentence . Traditional data augmentation algorithms for NMT Sennrich et al. (2016a); Zhang and Zong (2016); Zhou and Keung (2020); Nguyen et al. (2020) generally calculate cross-entropy losses on all data and use their summation to train the model:
However, this loss function is inaccurate for NAT due to the increase of data complexity. Sequence-level knowledge distillation works well on NAT by reducing the complexity of target data Zhou et al. (2020). In comparison, the target data generated by diverse distillation is relatively more complex. If NAT learns from the references indiscriminately, it will not eventually converge to any one reference but generate a mixture of all references.
Using the multi-reference dataset, we propose to train NAT with reference selection to evaluate the model output with better accuracy. We compare the model output with all reference sentences and select the one with the maximum probability assigned by the model. We train the model with only the selected reference:
In this way, we do not fit the model to all references but only encourage it to generate the nearest reference, which is an easier but more suitable objective for the model. Besides, when the ability of autoregressive teacher is limited, the NAT model can learn to ignore bad references in the data and select the clean reference for the training, which makes the capability of NAT not limited by a specific autoregressive teacher.
In addition to minimizing all losses or the selected loss , there is also an intermediate choice to assign different weights to reference sentences. We can optimize the log-likelihood of generating any reference sentence as follows:
The gradient of Equation 9 is equivalent to assigning weight to the cross-entropy loss of each reference sentence . In this way, the model focuses more on suitable references but also assigns non-zero weights to other references.
We use a linear annealing schedule with two stages to train the NAT model. In the first stage, we begin with the summation and linearly anneal the loss to . Similarly, we linearly switch to the selected loss in the second stage. We use and to denote the current time step and total training steps respectively, and use a constant to denote the length of the first stage. The loss function is:
where and are defined as:
In this way, the model learns from all references indiscriminately at the beginning, which serves as a pretraining stage that provides comprehensive knowledge to the model. As the training progresses, the model focuses more on the selected reference, which provides accurate training signals and gradually finetunes the model to the optimal state.
2.2 Efficient Calculation with CTC
To calculate the probability , the vanilla-NAT must set the decoder length to the length of . Therefore, calculating the probability of reference sentences requires running the decoder for at most times, which will greatly increase the training cost. Fortunately, for CTC-based NAT, the training cost is nearly the same since its decoder length is only determined by the source sentence. We only need to run the model once and calculate the probabilities of the reference sentences with dynamic programming, which has a minor cost compared with forward and backward propagations. In Table 1, we show the calculation cost of and for different models. We use CTC as the baseline model due to its superior performance and training efficiency.
2.3 Max-Reward Reinforcement Learning
Following Shao et al. (2019, 2021), we finetune the NAT model with the reinforcement learning objective Williams (1992); Ranzato et al. (2015):
where is the reward function and will be discussed later. The usual practice is to sample a sentence from the distribution to estimate the above equation. For CTC based NAT, cannot be directly sampled, so we sample from the equivalent distribution instead. We recover the target sentence by the collapsing function and calculate its probability with dynamic programming to estimate the following equation:
The reward function is usually evaluation metrics for machine translation (e.g., BLEU, GLEU), which evaluate the prediction by comparing it with the reference sentence. We use to denote the reward of prediction when is the reference. As we have references , we define our reward function to be the maximum reward:
By optimizing the maximum reward, we encourage the model to move towards the selected reference, which is the closest to the model. Otherwise, rewards provided by other references may mislead the model to generate a mixture of all references.
Experiments
Datasets We conduct experiments on major benchmark datasets for NAT: WMT14 EnglishGerman (EnDe, 4.5M sentence pairs) and WMT16 EnglishRomanian (EnRo, 0.6M sentence pairs). We also evaluate our approach on a large-scale dataset WMT14 EnglishFrench (EnFr, 23.7M sentence pairs) and a small-scale dataset IWSLT14 GermanEnglish (DeEn, 160K sentence pairs). The datasets are tokenized into subword units using a joint BPE model Sennrich et al. (2016b). We use BLEU Papineni et al. (2002) to evaluate the translation quality.
Hyperparameters We use 3 teachers for diverse distillation and set the seed to when training the -th teacher. We set the first stage length to . We use sentence-level BLEU as the reward. We adopt Transformer-base Vaswani et al. (2017) as our autoregressive baseline as well as the teacher model. The NAT model shares the same architecture as Transformer-base. We uniformly copy encoder outputs to construct decoder inputs, where the length of decoder inputs is as long as the source length. All models are optimized with Adam Kingma and Ba (2014) with and , and each batch contains approximately 32K source words. On WMT14 EnDe and WMT14 EnFr, we train AT for 150K steps and train NAT for 300K steps with dropout 0.2. On WMT16 EnRo and IWSLT14 DeEn, we train AT for 18K steps and train NAT for 150K steps with dropout 0.3. We finetune NAT for 3K steps. The learning rate warms up to within 10K steps in pretraining and warms up to within 500 steps in RL fine-tuning, and then decays with the inverse square-root schedule. We average the last 5 checkpoints to obtain the final model. We use GeForce RTX GPU for the training and inference. We implement our models based on fairseq Ott et al. (2019).
Knowledge Distillation For baseline NAT models, we follow previous works on NAT to apply sequence-level knowledge distillation Kim and Rush (2016) to make the target more deterministic. Our method applies diverse distillation with by default, that is, we use SeedDiv to generate 3 reference sentences for each source sentence.
Beam Search Decoding For autoregressive models, we use beam search with beam width 5 for the inference. For NAT, the most straightforward way is to generate the sequence with the highest probability at each position. Furthermore, CTC-based models also support beam search decoding optionally combined with n-gram language models Kasner et al. (2020). Following Gu and Kong (2020), we use beam width 20 combined with a 4-gram language model to search the target sentence, which can be implemented efficiently in C++https://github.com/parlance/ctcdecode.
2 Main Results
We compare the performance of DDRS and existing methods in Table 2. Compared with the competitive CTC baseline, DDRS achieves a strong improvement of more than 1.5 BLEU on average, demonstrating the effectiveness of diverse distillation and reference selection. Compared with existing methods, DDRS beats the state-of-the-art for one-pass NAT on all benchmarks and beats the autoregressive Transformer on most benchmarks with 14.7 speedup over it. The performance of DDRS is further boosted by beam search and 4-gram language model, which even outperforms all iterative NAT models with only one-pass decoding. Notably, on WMT16 EnRo, our method improves state-of-the-art performance levels for NAT by over 1 BLEU. Compared with autoregressive models, our method outperforms the Transformer with knowledge distillation, and meanwhile maintains 5.0 speedup over it.
We further explore the capability of DDRS with a larger model size and stronger teacher models. We use the big version of Transformer for distillation, and also add 3 right-to-left (R2L) teachers to enrich the references. We respectively use Transformer-base and Transformer-big as the NAT architecture and report the performance of DDRS in Table 3. Surprisingly, the performance of DDRS can be further greatly boosted by using a larger model size and stronger teachers. DDRS-big with beam search achieves 29.82 BLEU on WMT14 En-De, which is close to the state-of-the-art performance of autoregressive models on this competitive dataset and improves the state-of-the-art performance for NAT by over 1 BLEU with only one-pass decoding.
We also evaluate our approach on a large-scale dataset WMT14 En-Fr and a small-scale dataset IWSLT14 De-En. Table 4 shows that DDRS still achieves considerable improvements over the CTC baseline and DDRS with beam search can outperform the autoregressive Transformer.
3 Ablation Study
In Table 5, we conduct an ablation study to analyze the effect of techniques used in DDRS. First, we separately use the loss functions defined in Equation 7, Equation 8 and Equation 9 to train the model. The summation loss has a similar performance to the CTC baseline, showing that simply using multiple references is not helpful for NAT due to the increase of data complexity. The other two losses and achieve considerable improvements to the CTC baseline, demonstrating the effectiveness of reference selection.
Then we use different to verify the effect of the annealing schedule. With the annealing schedule, the loss is a combination of the three losses but performs better than each of them. Though the summation loss does not perform well when used separately, it can play the role of pretraining and improve the final performance. When is , the annealing schedule performs the best and improves by about 0.3 BLEU.
Finally, we verify the effect of the reward function during the fine-tuning. When choosing a random reference to calculate the reward, the fine-tuning barely brings improvement to the model. The average reward is better than the random reward, and the maximum reward provided by the selected reference performs the best.
4 DDRS on Autoregressive Transformer
Though DDRS is proposed to alleviate the multi-modality problem for NAT, it can also be applied to autoregressive models. In Table 6, we report the performance of the autoregressive Transformer when trained by the proposed DDRS losses. In contrast to NAT, AT prefers the summation loss , and the other two losses based on reference selection even degrade the AT performance.
It is within our expectation that AT models do not benefit much from reference selection. NAT generates the whole sentence simultaneously without any prior knowledge about the reference sentence, so the reference may not fit the NAT output well, in which case DDRS is helpful by selecting an appropriate reference for the training. In comparison, AT models generally apply the teacher forcing algorithm Williams and Zipser (1989) for the training, which feeds the golden context to guide the generation of the reference sentence. With teacher forcing, AT models do not suffer much from the multi-modality problem and therefore do not need reference selection. Besides, as shown in Table 1, another disadvantage is that the training cost of DDRS is nearly k times as large, so we do not recommend applying DDRS on AT.
5 Effect of Diverse Distillation
In the diverse distillation part of DDRS, we apply SeedDiv to generate multiple references. There are also other diverse translation techniques that can be used for diverse distillation. In this section, we evaluate the effect of diverse distillation techniques on the performance of DDRS. Besides SeedDiv, we also use HardMoe Shen et al. (2019) and Concrete Dropout Wu et al. (2020) to generate multiple references, and report their performance in Table 7. When applying other techniques for diverse distillation, the performance of DDRS significantly decreases. The performance degradation indicates the importance of high reference BLEU in diverse distillation, as the NAT student directly learns from the generated references.
6 Effect of Reward
There are many automatic metrics to evaluate the translation quality. To measure the effect of reward, we respectively use different automatic metrics as reward for RL, which include traditional metrics (BLEU Papineni et al. (2002), METEOR Banerjee and Lavie (2005), GLEU Wu et al. (2016)) and pretraining-based metrics (BERTScore Zhang* et al. (2020), BLEURT Sellam et al. (2020)). We report the results in Table 8. Comparing the three traditional metrics, we can see that there is no significant difference in their performance. The two pretraining-based metrics only perform slightly better than traditional metrics. Considering the performance and computational cost, we use the traditional metric BLEU as the reward.
7 Number of References
In this section, we evaluate how the number of references affects the DDRS performance. We set the number of references to different values and train the CTC model with reference selection. We report the performance of DDRS with different in Table 9. The improvement brought by increasing is considerable when is small, but it soon becomes marginal. Therefore, it is reasonable to use a middle number of references like to balance the distillation cost and performance.
8 Time Cost
The cost of preparing the training data is larger for DDRS since it requires training teacher models and using each model to decode the training set. We argue that the cost is acceptable since the distillation cost is minor compared to the training cost of NAT, and we can reduce the training cost to make up for it. In Table 10, we report the performance and time cost of models with different batch sizes on the test set of WMT14 En-De. DDRS makes up for the larger distillation cost by using a smaller training batch, which has a similar cost to the CTC model of 64K batch and achieves superior performance compared to models of 128K batch.
Related Work
Gu et al. (2018) proposes non-autoregressive translation to reduce the translation latency, which suffers from the multi-modality problem. A line of work introduces latent variables to model the nondeterminism in the translation process, where latent variables are based on fertilities Gu et al. (2018), vector quantization Kaiser et al. (2018); Roy et al. (2018); Bao et al. (2021) and variational inference Ma et al. (2019); Shu et al. (2020). Another branch of work proposes training objectives that are less influenced by the multi-modality problem to train NAT models Wang et al. (2019); Shao et al. (2019, 2020, 2021); Sun et al. (2019); Ghazvininejad et al. (2020); Shan et al. (2021); Du et al. (2021). Some researchers consider transferring the knowledge from autoregressive models to NAT Li et al. (2019); Wei et al. (2019); Guo et al. (2020a); Zhou et al. (2020); Sun and Yang (2020). Besides, some work propose iterative NAT models that refine the model outputs with multi-pass iterative decoding Lee et al. (2018); Gu et al. (2019); Ghazvininejad et al. (2019); Ran et al. (2020); Kasai et al. (2020). Our work is most related to CTC-based NAT models Graves et al. (2006); Libovický and Helcl (2018); Kasner et al. (2020); Saharia et al. (2020); Zheng et al. (2021); Gu and Kong (2020), which apply the CTC loss to model latent alignments for NAT. In autoregressive models, translations different from the reference can be evaluated with reinforcement learning Ranzato et al. (2015); Norouzi et al. (2016), probabilistic n-gram matching Shao et al. (2018), or an evaluation module Feng et al. (2020).
Our work is also related to the task of diverse machine translation. Li et al. (2016); Vijayakumar et al. (2018) adjust the beam search algorithm by introducing regularization terms to encourage generating diverse outputs. He et al. (2018); Shen et al. (2019) introduce latent variables with the mixture of experts method and use different latent variables to generate diverse translations. Sun et al. (2020) generates diverse translations by sampling different attention heads. Wu et al. (2020) train the translation model with concrete dropout and samples different models from a posterior distribution. Li et al. (2021) generate different translations for the input sentence by mixing it with different sentence pairs sampled from the training corpus. Nguyen et al. (2020) augment the training set by translating the source-side and target-side data with multiple translation models, but they do not evaluate the diversity of the augmented data.
Conclusion
In this paper, we propose diverse distillation with reference selection (DDRS) for NAT. Diverse distillation generates a dataset containing multiple references for each source sentence, and reference selection finds the best reference for the training. DDRS demonstrates its effectiveness on various benchmarks, setting new state-of-the-art performance levels for non-autoregressive translation.
Acknowledgement
We thank the anonymous reviewers for their insightful comments. This work was supported by National Key R&D Program of China (NO.2017YFE9132900).