Noisy Channel Language Model Prompting for Few-Shot Text Classification
Sewon Min, Mike Lewis, Hannaneh Hajishirzi, Luke Zettlemoyer
Introduction
Prompting large language models, by prepending natural language text or continuous vectors (called prompts) to the input, has shown to be promising in few-shot learning (Brown et al., 2020). Prior work has proposed methods for finding better prompt (Shin et al., 2020; Li and Liang, 2021; Lester et al., 2021) or better scoring of the output from the model (Zhao et al., 2021; Holtzman et al., 2021). These studies directly predict target tokens to determine the prediction for an end task. Despite promising results, they can be unstable with high variance across different verbalizers (text expression for labels) and seeds, and the worst-case performance is often close to random (Perez et al., 2021; Lu et al., 2021).
In this paper, we introduce alternative channel models for prompted few-shot text classification with large language models, inspired by noisy channel models in machine translation (Brown et al., 1993; Koehn et al., 2003; Yu et al., 2017; Yee et al., 2019) and their extensions to other tasks (Yogatama et al., 2017; Lewis and Fan, 2018). Unlike direct models that compute the conditional probability of the label token given the input, channel models compute the conditional probability of the input given the output (Figure 1). Intuitively, channel models are required to explain every word in the input, potentially amplifying training signals in the low data regime. We study the impact of channel models for language model prompting where the parameters of the language model are frozen. In particular, we compare channel models with their direct counterparts for (1) demonstration methods, either concatenation-based (Brown et al., 2020) or our proposed, ensemble-based (Section 4.1.3), and (2) prompt tuning (Lester et al., 2021).
Our experiments on eleven text classification datasets show that channel models outperform their direct counterparts by a large margin. We attribute the strong performance of channel models to their stability: they have lower variance and significantly higher worst-case accuracy then their direct counterparts over different verbalizers and seeds. We additionally find a direct model with head tuning—tuning the LM head while freezing other parameters—is surprisingly effective, often outperforming direct models with other forms of tuning. While different methods are preferred given different conditions, the channel model with prompt tuning (denoted as channel prompt tuning) significantly outperforms all direct baselines when (1) the training data is imbalanced, or (2) generalization to unseen labels is required.
In summary, our contributions are three-fold:
We introduce a noisy channel approach for language model prompting in few-shot text classification, showing that they significantly outperform their direct counterparts for both demonstration methods and prompt tuning.
We find particularly strong performance of channel models over direct models when the training data is imbalanced or generalization to unseen labels is required.
Based on extensive ablations, we provide recommendations between different models (direct vs. channel and prompt tuning vs. head tuning) based on given conditions such as the target task, the size of training data, the number of classes, the balance between labels in the training data, and whether generalization to unseen labels is required.
Related Work
Let and be the input and the output, respectively. The most widely used models, denoted as direct models, compute . In contrast, noisy channel models maximize (Shannon, 1948; Brown et al., 1993). We follow Yu et al. (2017); Yee et al. (2019) in using the terms direct models and channel models. They are often referred as discriminative models and generative models in prior work (Yogatama et al., 2017; Lewis and Fan, 2018). In principle, these two distinctions are not always equivalent, e.g., a model that computes is generative but not a channel model. While the noisy channel approach has been the most successful in machine translation (Yamada and Knight, 2001; Koehn et al., 2003; Yu et al., 2017; Yee et al., 2019), it has also been studied in more general NLP tasks. Prior work provides a theoretical analysis that channel models approach their asymptotic errors more rapidly than their direct counterparts (Ng and Jordan, 2002), and empirically shows that channel models are more robust to distribution shift in text classification (Yogatama et al., 2017) or question answering (Lewis and Fan, 2018), and in a few-shot setup (Ding and Gimpel, 2019).
In this paper, we explore channel models using a large language model on a wide range of text classification tasks, focusing on prompt-based few-shot learning.
2 Few-shot Learning
Prior work in few-shot learning has used different approaches, including semi-supervised learning with data augmentation or consistency training (Miyato et al., 2017; Clark et al., 2018; Xie et al., 2020; Chen et al., 2020) and meta learning (Finn et al., 2017; Huang et al., 2018; Bansal et al., 2020). Recent work has introduced prompting (or priming) of a large language model. For example, Brown et al. (2020) proposes to use a concatenation of training examples as a demonstration, so that when it is prepended to the input and is fed to the model, the model returns the output following the pattern in the training examples. This is especially attractive as it eliminates the need for updating parameters of the language model, which is often expensive and impractical. Subsequent work proposes alternative ways of scoring labels through better model calibration (Zhao et al., 2021; Holtzman et al., 2021), or learning better prompts, either in a discrete space (Shin et al., 2020; Jiang et al., 2020; Gao et al., 2021) or in a continuous space (Li and Liang, 2021; Lester et al., 2021; Liu et al., 2021; Zhong et al., 2021; Qin and Eisner, 2021). Almost all of them are direct models, computing the likelihood of given with the prompts.
Our work is closely related to two recent papers. Tam et al. (2021) studies a label-conditioning objective for masked language models; although this is not strictly a generative channel model, conditioning on the output is similar to our work. However, they are still optimizing a discriminative objective, and inference at test time is the same as with the direct model. Holtzman et al. (2021) explores zero-shot models that compute the probability of given based on Pointwise Mutual Information, but with a restriction that the input and the output are interchangeable. To the best of our knowledge, our work is the first that uses a noisy channel model for few-shot language model prompting for classification, and also the first to draw the connection with the noisy channel literature.
Formulation
We focus on text classification tasks. The goal is to learn a task function , where is the set of all natural language texts and is a set of labels. We consider three formulations.
Direct computes distributions of labels given the input : . This is the most widely used method in modern neural networks.
Direct++ is a stronger direct model that computes instead of , following the method from Holtzman et al. (2021) and the non-parametric method from Zhao et al. (2021). This approach is motivated by the fact that language models can be poorly calibrated and suffer from competition between different strings with the same meaning. This approach is used for the demonstration methods in Section 4.1.
Method
When learning a task function , we also assume a pre-defined verbalizer which maps each label into a natural language expression. For example, if the task is sentiment analysis with , an example input text would be “A three-hour cinema master class” and an example would have “It was great” and “It was terrible”. In a few-shot setup, we are also given a set of training examples .
We are interested in methods where there are no trainable parameters (Section 4.1) or the number of trainable parameters is very small, typically less than 0.01% of the total (Section 4.2). This follows prior observations that updating and saving a large number of parameters for every task is expensive and often infeasible (Rebuffi et al., 2017; Houlsby et al., 2019; Lester et al., 2021).
In demonstration methods, there are no trainable parameters. We explore three ways of making a prediction, as summarized in Table 1.
1.2 Concat-based demonstrations
1.3 Ensemble-based demonstrations
2 Tuning methods
We also explore methods that tune a very limited number of model parameters, as summarized in Figure 2. We study head tuning (Section 4.2.1) and transformation tuning (Section 4.2.2) for direct models. We also consider prompt tuning (Section 4.2.3) for both direct and channel models, which we refer as direct prompt tuning and channel prompt tuning, respectively. All models share the same input-output interface with the zero-shot setup in Table 1 during training and inference.
2.2 Transformation tuning
2.3 Prompt tuning
Experimental Setup
We report results for eleven text classification datasets, following Zhang et al. (2015) and Gao et al. (2021): SST-2 (Socher et al., 2013), SST-5 (Socher et al., 2013), MR (Pang and Lee, 2005), CR (Hu and Liu, 2004), Amazon (McAuley and Leskovec, 2013), Yelp (Zhang et al., 2015), TREC (Voorhees and Tice, 2000), AGNews (Zhang et al., 2015), Yahoo (Zhang et al., 2015), DBPedia (Lehmann et al., 2015) and Subj (Pang and Lee, 2004). The datasets include a varied number of classes per task, from 2 to 14. See Table 10 in Appendix A for dataset samples.
2 Training Data
We follow all the hyperameters and details from prior work (Appendix B) which eliminates the need for a held-out validation set. The very limited data is better used for training rather than validation, and cross-validation is less helpful when the validation set is extremely small (Perez et al., 2021).
3 Language Models
We use GPT-2 (Radford et al., 2019) for the LM. We primarily use GPT-2 Large but also experiment with varying sizes (Small, Medium, Large and X-Large) for the ablations in Appendix C. While we only experiment with GPT-2, our experiments are easily extendable to other causal language models.
4 Evaluation
We use accuracy as a metric for all datasets.
We experiment with 4 different verbalizers (taken from Gao et al. (2021); full list provided in Appendix A), 5 different random seeds for sampling training data, and 4 different random seeds for training. We then report Average accuracy and Worst-case accuracy.We also report standard deviation and best-case accuracy in the Appendix. We consider the worst-case accuracy to be as important as the average accuracy given significantly high variance of few-shot learning models, as shown in previous work (Zhao et al., 2021; Perez et al., 2021). The worst-case accuracy is likely of more interest in high-risk applications (Asri et al., 2016; Guo et al., 2017).
Other implementation details are in Appendix B. All experiments are reproducible from github.com/shmsw25/Channel-LM-Prompting.
Experimental Results
This section reports results from demonstration methods (Section 6.1), tuning methods (Section 6.2) and ablations (Section 6.3). Discussion is provided in Section 7.
Table 3 shows the performance of demonstration methods.
Direct++ significantly outperforms the naive direct model across all setups, indicating that using instead of is highly beneficial as claimed by Holtzman et al. (2021); Zhao et al. (2021).
Our proposed, ensemble-based method is better than the concat-based method in direct models, by 7% absolute in the average accuracy and the worst-case accuracy, when macro-averaged across all datasets.
In contrast, the ensemble-based method is not always better in channel models; it is better only on the datasets with long inputs. We conjecture that the ensemble-based method may suffer when labels in the training data are not balanced, which direct++ explicitly takes into account as described in Zhao et al. (2021).
In a few-shot setting, channel models outperform direct models in almost all cases. The strongest channel model outperforms the strongest direct model by 3.1% and 7.2% absolute, in terms of the average accuracy and the worst-case accuracy, respectively.
Standard deviation and the best-case accuracy are reported in Table 11 and Table 12 in the Appendix. They indicate strong performance of channel models can be attributed to their low variance. The highest best-case accuracy is achieved by direct++ on most datasets, but it has a higher variance, having lower average and the worst-case accuracy than channel models.
Performance of direct models sometimes degrades in a few-shot setting, which is also observed by prior work (Zhao et al., 2021). This is likely because demonstrations provided by the training data may cause the model to be miscalibrated and easily biased by the choice of demonstrations. However, channel models achieve few-shot performance that is significantly better than zero-shot methods on all datasets.
2 Main Results: Tuning Methods
Table 4 shows the performance of tuning methods.
When using prompt tuning, channel models consistently outperform direct models by a large margin on all datasets. Improvements are 13.3% and 23.5% absolute in the average and the worst-case accuracy, respectively.
Standard deviation and the best-case accuracy are reported in Table 13 in the Appendix. Consistent with the findings in Section 6.1, the strong performance of channel prompt tuning can be explained by the low variance of channel prompt tuning. Direct prompt tuning often achieves higher best-case accuracy; however, due to its high variance, its overall accuracy is lower, with significantly lower worst-case accuracy.
We find that head tuning is a very strong method, despite often being omitted as a baseline in prior work. It significantly outperforms direct prompt tuning in all cases. It also outperforms channel prompt tuning on some datasets, particularly significantly on TREC and Subj. For these datasets, the task—finding the type of the answer to the question or identifying the subjectivity of the statement—is inherently different from language modeling, and likely benefits from directly updating the LM parameters, rather than using the LM as a black box.
Still, channel prompt tuning outperforms direct head tuning on most datasets. The largest gains are achieved on Yahoo and DBPedia. In fact, on these datasets, channel prompt tuning even outperforms all finetuning—finetuning all parameters of the LM—which achieves 48.9/43.8 on Yahoo and 66.3/50.4 on DBPedia. We conjecture that using on these datasets naturally requires generalization to unseen labels due to the large number of classes ( and ), where channel prompt tuning significantly outperforms direct models, as we show in Section 6.4.
3 Ablations
For the ablations, we report experiments on SST-2, MR, TREC and AGNews, using one train seed (instead of four), and four verbalizers and five data seeds (as in main experiments).
On binary datasets (SST-2 and MR), we vary the label imbalance in the training data with . Specifically, let and , i.e., the ratio of in the training data. We vary to be . means the labels are perfectly balanced, and means that labels in the training data only include . We additionally compare with upsampling baselines where we upsample training examples with infrequent labels so that the model has seen an equal number of examples per label during training.
Results are reported in Figure 4. All direct models are sensitive to the imbalance in training data, even though they benefit from upsampling when is small. Channel prompt tuning is insensitive to the imbalance, and significantly outperforms direct models when is small; it even outperforms all finetuning when . When is near to 0.5, direct head tuning matches or outperforms channel prompt tuning.
It is also worth noting that direct prompt tuning with upsampling matches or outperforms all finetuning and head tuning when is small.
4 Generalization to unseen labels
We experiment with a challenging scenario where the model must generalize to unseen labels. While it may be seen as an extreme scenario, this is often a practical setting, e.g., the problem is defined with a set of labels but later an addition of the new label may be needed.
First, we sample training examples as in main experiments but excluding one random label, so that at least one label at test time was unseen during training. Table 5 reports the results. All direct models are unable to predict the label that is unseen at training time. However, channel prompt tuning can predict unseen labels and achieves considerably better performance than zero-shot. It outperforms all finetuning on 2-way classification datasets, and outperforms head tuning on five datasets except for TREC on which head tuning achieves very strong performance on seen labels.
Next, we run zero-shot transfer learning, where the model is trained on one dataset and is tested on another dataset. Here, head tuning is not applicable when the labels are not shared between two datasets. Figure 5 shows the results. Channel prompt tuning outperforms all direct models including all finetuning on all datasets except for TREC. It is particularly competitive when the tasks are inherently similar, e.g., transfer between 2-way sentiment analysis and 5-way sentiment analysis in the first three figures. In fact, in such cases, performance is close to the models trained on in-domain data. When tasks are inherently different, e.g., the rest of the figures in Figure 5, gains over zero-shot performance are relatively small; we think more work should be done to make cross-task transfer better and to discover when it is possible.
Discussion & Conclusion
In this work, we introduced a noisy channel approach for few-shot text classification through LM prompting, where we either provide demonstrations to the LM or tune the prompt embeddings in the continuous space. Our experiments on eleven datasets show that channel models significantly outperform their direct counterparts, mainly because of their stability, i.e., lower variance and better worst-case accuracy. We also found that direct head tuning is more competitive than previously thought, and different methods are preferred given different conditions. Specifically, channel prompt tuning is preferred in the following scenarios.
Channel prompt tuning is more competitive when there are fewer training examples. We hypothesize two reasons: (1) Channel models are more stable (i.e., achieve low variance and high worst-case accuracy), unlike direct models that are highly unstable with small (Zhao et al., 2021; Perez et al., 2021; Lu et al., 2021). (2) Channel models provide more signals by requiring the model to explain the input word-by-word (as claimed in Lewis and Fan (2018)) which is beneficial in the low data regime.
When the training data is even slightly imbalanced, no direct models are competitive. We think this is because the LM head relies too much on unconditional distributions of labels. Channel prompt tuning is less sensitive because labels are only a conditioning variable. Label imbalance in the training data is a real-world problem, especially when is small and is large. We thus suggest this is an important area for future work.
All direct models are unable to predict labels that are unseen during training, indicating that they overfit in the label space. In contrast, channel models can predict unseen labels, likely because the label space is indirectly modeled. This is in line with prior work that shows channel models are more competitive under a distribution shift (Yogatama et al., 2017; Lewis and Fan, 2018).
If the task is too different from language modeling even with carefully chosen verbalizers (e.g., TREC and Subj), head tuning outperforms prompt tuning. This is likely because it benefits from directly updating the parameters of the LM. This may mean that causal LMs are not suitable for all tasks, or we need more sophisticated methods to apply causal LMs for such tasks without updating the LM parameters.
While we show that channel models are competitive in few-shot text classification, there are limitations that provide avenues for future work. First, it is not as easy to use channel models for non classification tasks where modeling prior distributions is non-trivial. We think future work can obtain the prior with a separate model and incorporate it to the conditional LM as done by Lewis and Fan (2018), potentially with beam search decoding as in Yu et al. (2017); Yee et al. (2019).
Second, while this paper focuses on causal LMs, it is an open question how to use a channel model with masked LMs. Although we think channel models are not inherently restricted to causal LMs, the specific way in which existing masked LMs are pretrained makes it hard to use channel models without updating the LM parameters, e.g., masked LMs are not trained to generate long sentences. One recent approach uses a label-conditioning objective (Tam et al., 2021) as a clever way to introduce a channel-like model with existing masked LMs. Extending and further integrating these different approaches would be important for using channel models in a wider range of scenarios.
Acknowledgements
We thank Ari Holtzman, Eric Wallace, Gabriel Ilharco, Jungsoo Park, Myle Ott, Peter West and Ves Stoyanov for their helpful comments and discussion. This research was supported by NSF IIS-2044660, ONR N00014-18-1-2826, an Allen Distinguished Investigator Award, and a Sloan Fellowship.
References
Appendix A Samples & Verbalizers
Table 10 shows samples from each dataset. Table 6 shows a list of verbalizers (four for each dataset), mainly taken from Gao et al. (2021) and label words included in the original data.
Appendix B Implementation Details
We use PyTorch (Paszke et al., 2019) and Huggingface Transformers (Wolf et al., 2020). For MR, we use the sentence polarity dataset version 1.0. We use the batch size of 32 and the sequence length of 128 for datasets with short input text (SST-2, SST-5, MR, TREC) and the batch size of 16 and the sequence length of 256 for datasets with long input text (AGNews, Amazon, Yelp, DBPedia, Yahoo, Subj). When the concat-based demonstration method is used, the sequence length is multiplied by the number of training examples, yet is bounded by 1024 which is a strict limit of GPT-2.
For all finetuning experiments, we train the model for 100 global steps. We use the loss divided by the number of all tokens in the batch. We use Adam optimizer (Kingma and Ba, 2015) with no weight decay and no warmup steps. For head tuning, transformation tuning and prompt tuning, we use the learning rate and choose the one that gives the lowest training loss on average in order to eliminate the need of the validation data. The chosen learning rate values are reported in Table 7. For all finetuning, we use the learning rate of . For prompt tuning, we use prompt tokens which embeddings are initialized from a random subset of the top vocabularies, following the original paper (Lester et al., 2021).
Appendix C Additional Results
Table 11, 12 and 13 report the average accuracy, the variance, the best-case accuracy and the worst-case accuracy using the concat-based demonstration, the ensemble-based demonstration and the tuning methods, respectively. Results consistently indicate that channel models achieve significantly lower variance and higher worst-case accuracy. The best-case accuracy is often achieved by direct models, but channel models outperform direct models on average.
We vary the size of LMs and report the average and the worst-case accuracy in Figure 6. The trends—no matter the best performance is achieved by channel prompt tuning or direct head tuning—are fairly consistent across varying size of LMs.