Conformal Language Modeling
Victor Quach, Adam Fisch, Tal Schuster, Adam Yala, Jae Ho Sohn, Tommi S. Jaakkola, Regina Barzilay
Introduction
Language models (LMs) have emerged as powerful tools for solving natural language processing (NLP) tasks. Given an input prompt, LMs generate a response from some predicted distribution over output text sequences. For modern models, these generations are often coherent and contextually relevant. At the same time, these generations can still contain mistakes, and lack certain aspects of robustness and reliability in terms of providing accurate, trustworthy predictions [27; 31; 38; 42; 63; 72]. Unfortunately, quantifying the uncertainty in LM outputs has remained a major challenge.
Conformal prediction is a popular model-agnostic and distribution-free method for creating prediction sets that contain the correct answers with high probability [2; 3; 4; 6; 36; 55; 69]. Applying conformal prediction to generative models such as LMs, however, is challenging due to (a) the unbounded nature of their output space (i.e., all possible text sequences), and (b) the limited available (tractable) mechanisms for exploring all possible predictions. In particular, LMs can typically only approximately search or sample candidate responses. Furthermore, while several possible responses might be acceptable (e.g., correct or factual), small differences can result in abrupt changes in coherence or meaning.
In this paper, we propose an extension of conformal prediction that is tailored specifically to generative LMs. We only assume that the (potentially black-box) LM that is given to us can be used to sample diverse output sequences, together with their evaluated model likelihoods (i.e., the output token sequence logits). Like conformal prediction, our method offers a rigorous coverage guarantee by constructing prediction sets that, in our case, provably contain at least one acceptable response with high probability. Unlike conformal prediction, however, we do not enumerate the entire output space (which is impossible). Instead, we derive a calibrated stopping rule for sampling different outputs from the LM that get added to a growing output set of candidates, until we are confident that the output set is sufficient. Since not all samples from the LM may be high quality (e.g., some may be redundant, incoherent, or have lower confidence), we also simultaneously calibrate a rejection rule for removing candidates from the output set—while still ensuring that our coverage bound is not violated. This gives the benefit of making our output sets not only accurate, but also precise (i.e., small).
Contributions. In summary, our main results are as follows:
We bridge the gap between conformal prediction and LMs by calibrating the sampling of output sets, rather than enumerating and selecting candidate responses directly from the output space;
We extend multi-label conformal prediction to identify confident components of long generations;
Though limitations apply, we demonstrate valid risk control on multiple diverse tasks with different LMs, while still retaining meaningful output sets that are precise on average compared to baselines.
Related work
Conformal prediction and risk control. Our work adds to the rich collection of tools for uncertainty estimation and risk control for machine learning algorithms [2; 3; 5; 6; 16; 18; 35; 36; 68; 70; 71, inter alia]. These techniques were previously extended and applied in the language domain to classification with finitely-many classes [14; 15; 27], to token-level predictions [11; 54], and to reliably accelerate LMs [34; 57; 59]. Here, we address the emerging challenge of providing reliable prediction sets for unbounded, free-text generation—which previous methods are unequipped for. The distribution-free, finite-sample performance guarantees that we derive are similar to those given by prediction sets or regression intervals in standard conformal prediction [2; 50; 69], but with slightly relaxed “correctness” criterions [9; 14]. In particular, we build on the groundwork set by Angelopoulos et al. , which provides a general methodology for calibrating any risk function that is controllable via some low-dimensional hyper-parameter configuration. We extend their framework to handle sampling-based algorithms that can effectively be used for LMs, and that, critically, do not require enumerating the full output space (which is intractable in our case). Most relevant to our work in LMs, other recent approaches have built on conformal principles to construct confidence intervals for generative diffusion models over images [23; 64]. These methods do not directly translate to LMs, however, as they only provide non-combinatorial confidence intervals on the pixel-level.
Uncertainty estimation in LMs. As the use of LMs in-the-wild quickly grows, there is increasing interest in obtaining and expressing meaningful confidence estimates for each output. Recent studies show that the logits of out-of-the-box LMs tend to exhibit overconfidence, even when wrong [10; 29; 43; 67]. Recent alignment techniques degrade this even further [29; 48]. Most current mitigation approaches focus on introducing linguistic cues [39; 80] or empirical post-hoc logit calibration [25; 29; 44; 78]. Such heuristics, however, don’t provide any concrete guarantees. In this work, we develop similar techniques to improve the output of the underlying LM. Our methods are model agnostic and provide rigorous guarantees. Our conformal component selection (§4.4) also relates to recent self-consistency work that builds on the empirical observation that repeated similar samples are more likely to be correct [45; 72], and cross-sample entailment can approximate uncertainty . Unlike previous work that uses a fixed number of re-samples and compares full outputs, we (1) introduce a dynamic stopping rule to reduce the number of samples, (2) extend this concept to semantically compare sub-components of long text outputs, and (3) conformalize the process to provide proper guarantees.
Reliable generation. It is common practice to post-hoc apply classifiers and filters on top of LM generations for various quality goals such as preventing toxicity [17; 53; 73], verifying grounding against sources [7; 41; 77], or re-ranking the set of decoded outputs . Our work provides a systematic and reliable approach for filtering or flagging poor-quality outputs—both at a full generation and component level—and can also readily incorporate additional signal from auxiliary classifiers. For example, we demonstrate in our experiments using off-the-shelf natural language inference (NLI) models [8; 30; 56; 65; 74; 79] to help guide the selection of individual, confident components in text summarization (i.e., sentences that are fully entailed by the larger text [13; 22; 33; 58]).
Background
We begin with a brief review of conformal prediction and general risk control (see also ). Here, and in the rest of the paper, upper-case letters () denote random variables; lower-case letters () denote constants, and script letters () denote sets, unless specified. All proofs are in Appendix B.
Let , be exchangeable random variables. Let random variable be the nonconformity score of , where is fixed. For , define the prediction (based on the first examples) at as
Note that the coverage property expressed in Theorem 3.1 is marginal over the draw of calibration and test data. The recent Learn Then Test (LTT) framework of Angelopoulos et al. extends conformal prediction to control the expectation of any loss function (conditional on the draw of calibration data) by reframing hyper-parameter selection as a statistical multiple hypothesis testing problem.
Conformal language modeling
We now introduce our method for generating uncertainty sets for LMs. At a high level, our procedure consists of three main steps to sample and return an collection of plausible output predictions:
Sample. A new candidate response is sampled from our language model.
Accept or reject. The sample is added to the growing output set, as long as it is diverse (e.g., maximum overlap with any other element is ) and confident (e.g., the LM likelihood is ).
Stop or repeat. Using a set-based scoring function, we check if the confidence in the current set is . If it is, then we stop and return the current set. Otherwise we return to Step 1.
is a configuration that we calibrate to find a valid setting, , that controls the risk of our output sets. In the following, we more carefully define our setting and notation (§4.1), and then describe our sampling (§4.2) and calibration algorithms (§4.3). Then, in §4.4, we provide an additional extension for highlighting confident generation components—i.e., subsections of our full generations that are independently likely to be correct, even if the full generation is not.
Let be an alphabet (a non-empty, finite set of tokens such as ) from which all possible output strings, , are composed, i.e. .We write to denote the Kleene closure of a set , i.e., . We assume that we are given a generative model that defines a conditional probability distribution given some input prompt (where may be text, or another modality such as an image), which we can sample from to obtain candidate output strings, . Following Fisch et al. , for every input prompt , we assume access to some “admission” function that is used to measure the acceptability of a given sample . Intuitively, tells us if an output is “good enough”.
Continuing our radiology report example from §1, is the input X-ray, is the report, and is our image-to-text LM (in English). Given and some “ground truth” report (e.g., written by a radiologist), might measure if and agree on all findings. Note that, in practice, it may be hard to exactly define such an , or at least an that is automatically computable without extensive manual annotation. In Appendix E we show that it is also sufficient to only require access to a conservative admission function, , where we have . For instance, might measure exact match on a word-for-word basis between and , instead of accounting for differences in dictation. We explore different tasks and admission functions in our experiments in §5.
2 Conformal sampling with rejection
3 Calibration with Learn Then Test
4 Conformal selection of individual components
Furthermore, by the union bound, Eq. (1) and Eq. (9) hold simultaneously with probability .
Experimental setup
In this section, we briefly describe our experimental setup. Appendix F contains additional details.
Radiology report generation. As motivated in §1, we apply our method to chest X-ray radiology report generation using the MIMIC-CXR dataset. For our LM, we fine-tune an encoder-decoder architecture based on a pretrained ViT image encoder and a GPT2-small text decoder. To judge admission, we use the popular Clinical Efficacy metric [40; 47] to check if the 14 labels predicted by an auxiliary CheXbert model on the generated report exactly match the labels predicted by the same CheXbert model for a reference report from a radiologist. Similarly, a component (here a sentence including a finding) is defined to be admissible if it has a ROUGE-L score (picked through empirical validation), when compared to any component directly extracted from the reference.
News summarization. We also apply our method to news article text summarization using the CNN/DM dataset. For our LM, we finetune a T5-XL model. We define a candidate generation to be admissible if it has a ROUGE-L score higher than , when compared to all available reference summaries from human annotators. Like MIMIX-CXR, we define a component to be admissible if it has a ROUGE-L score when compared to components extracted from human summaries.
Open-domain question answering. Finally, we apply our method to open-domain question answering using the TriviaQA dataset. For this task, we sample answers from the LLaMA-13B LM in the few-shot setting , without any additional fine-tuning. Since answers are limited to one or few tokens, a candidate output generation is acceptable only if it exactly matches an annotated reference answer (after minor normalization for removing articles, casing, and punctuation). Furthermore, since the expected answers are short and fairly atomic, we do not evaluate component-level confidence.
2 Scoring functions
As discussed in §4.2, our method is implementation-agnostic and can support different choices of quality function , similarity function , and set scoring function . For our purposes, we show that a straightforwar approach is to simply use transformations on the model likelihoods (from token logits). Specifically, we define using the likelihood function of the base LM, with length-normalization . We use ROUGE-L for . For , we experiment with the following variants:
First-K. As a baseline, we score a set by its size, , and do not use rejection. This corresponds to the number of samples taken, and follows the intuition from our toy example in §4.2.
Max. The scoring function stems from the intuition that a set is only as good as its best element, and defines .
Sum. Alternatively, we also use the sum of item-level scores: .
3 Metrics
Our main motivation is to produce valid confidence sets that are also precise. To reflect this, we measure both the loss of our sets (which is guaranteed to satisfy our prescribed limits), as well as both (a) the relative number of “excess” samples taken from our model (including rejected samples, see also Eq. (7)), and (b) the ultimate output size of the prediction set (after rejection). Both metrics are important, as over-sampling wastes computational budget (or expensive API calls), while large output sets can be unwieldy to use and overall less helpful as an uncertainty quantification tool. We measure results and compute the AUC over the range of achievable or (using a fixed ), excluding trivial values (e.g., that a policy that always returns the first generation would satisfy).
Experimental results
We now present our main results. In all plots, solid lines give the mean over 100 trials and shaded regions show the standard deviation. Additional experimental results are reported in Appendix G.
Validity of conformal sampling with rejection. As per Theorem 4.2, we observe in Figure 2 that our conformal sampling approach is valid, as the average set loss often matches but never exceeds the target risk level. Methods that have access to the model logits (namely Max and Sum) are close to the diagonal line, indicating that they are not overly conservative.
Prediction efficiency. The likelihood-based approaches outperform the uniform First-K baseline across all three tasks. For example, as Figure 2(c) shows, the AUC of expected set size of Max and Sum are both less than half the AUC of First-K in the QA task. In tasks with longer output texts, First-K produces competitive set sizes across all achievable . However, it is being overly conservative on easy examples at the expense of hard ones. This is revealed when plotting the relative number of excess samples, where the Max scoring function largely outperforms Sum and First-K.
Individual components. We evaluate two scoring functions to apply our conformal component selection method. A first method Span-logits extracts the likelihood of a candidate component, as produced by the language model. Since those likelihoods are conditioned on previous context, this scoring function may underestimate the score of a correct component if it follows an incorrect component. Instead, we use an application-specific Classifier to assign a conformity score to each component. We compare these methods to a Random baseline which attributes a random score to any pair. Figure 3 shows that by modeling components independently, we produce more effective (larger) sets. We include qualitative results in Appendix H.
Conclusion
Reliably using language models (LMs) in real-world tasks inevitably requires meaningful uncertainty quantification. In this paper, we introduced a novel approach to conformal prediction that allows a user to sample prediction sets from generative LMs with infinite, combinatorial output spaces, while retaining desirable statistical guarantees. Our method bridges the gap between standard conformal prediction and LM inference techniques by calibrating a stopping rule for an algorithm that iteratively grows an output prediction set by sampling new generations (with rejection). Moreover, we provide a method for separately identifying answer components that we are more confident are accurate. This can help users better understand the quality of LM answers, including which parts may be incorrect (and vice versa) within a larger, verbose output. Finally, we demonstrate our method on three popular LM applications. Compared to common practices (e.g., first ), we obtain more efficient prediction sets, both in terms of size and samples required, leading to more effective outputs that are also less costly to obtain.
References
Appendix A Limitations and broader impact
Appendix B Proofs
B.2 Proof of Theorem 4.2
Therefore, Equation 1 holds for any . In particular, it holds for
B.3 Proof of Proposition 4.4
Appendix C Effects of truncated sampling
Appendix D Pareto testing
Appendix E Admission functions
Note that, as briefly discussed in §4.1, is also valid if a conservative, “approximate” admission function is used in place of a “complete” during calibration.
The following proof is analogous to that of Propostion 4.4. Let
Applying Theorem 4.2 gives that the left hand side is w.p. . ∎
Appendix F Additional experimental details
In this section, we provide additional details regarding the experiments conducted for the three tasks discussed in Section 5. Our code will be released after the review process.
For the radiology report generation experiment, we utilized the labeled MIMIC-CXR and MIMIC-CXR-JPG datasets [Johnson et al., 2019]. The MIMIC-CXR dataset can be accessed at https://physionet.org/content/mimic-cxr/2.0.0/ under the PhysioNet Credentialed Health Data License 1.5.0. Similarly, the MIMIC-CXR-JPG dataset is available at https://physionet.org/content/mimic-cxr-jpg/2.0.0/ under the same license.
We start with the standard splits prescribed in MIMIC-CXR-JPG. However, we further divide the training set into a train set and a dev set using a 0.9/0.1 ratio. The train set is used for training the model, using the validation set for early stopping. We then exclusively use the dev set for conformal prediction experiments. Subsequently, we filtered the dataset to include only anterior to posterior (AP) or posterior to anterior (PA) views and retained only one image per report. Furthermore, we removed examples where the report did not start with the phrase “FINAL REPORT” as these reports often contained a summary of the findings at the beginning, inadvertently leaking the answer we aimed to generate with the model. Table F.1 provides a statistical overview of the resulting dataset.
Each image was resized and cropped to a resolution of 224x224. Following prior methodology [Miura et al., 2021], we split each report into a prompt part and a findings part (which may also contain the impressions section) by identifying one of the following phrases: “FINDINGS AND IMPRESSION”, “FINDINGS” or “IMPRESSION”.
Model
The image encoder used in our experiment was a Vision Transformer (ViT) model pretrained on ImageNet-21k at a resolution of 224x224. Specifically, we utilized the google/vit-base-patch16-224-in21k model available in the Transformers library Wolf et al. . The text decoder was a GPT2-small model (gpt2 on HuggingFace). We trained the model with a batch size of 128 distributed over 8 GPUs, resulting in a batch size of 16 per GPU. The AdamW optimizer was employed with , , and . The learning rate was set to . The training process consisted of 10 epochs, and the total training time on 8 RTX A6000 GPUs was approximately 11 hours.
Generations
Candidate reports were sampled from the model using default arguments from the Transformers library, i.e. , and temperature = 1. Each generated report is then evaluated using a trained CheXbert model Smit et al. . The CheXbert model is available at https://stanfordmedicine.box.com/ under the Stanford Academic Software License. The CheXbert model labels each report for 14 conditions, assigning one of the following labels: “Blank,” “Positive,” “Negative,” or “Uncertain.”
To determine the admission of a candidate report, we compare it with a reference (human) report from the MIMIC dataset. If the candidate report matches all 14 labels of the reference report, the admission function returns 1; otherwise, it returns 0.
Components
We define a component as a sentence delimited by a period. The component-level admission function is defined based on how well a sentence“almost matches” one of the reference sentences. Two sentences are considered to “almost match” if their ROUGE score is above 0.4. If a sentence almost matches a reference sentence, the component-level admission function returns 1; otherwise, it returns 0.
F.2 Open-domain question answering
We use the TriviaQA [Joshi et al., 2017] dataset available at https://nlp.cs.washington.edu/triviaqa/ under the Apache License Version 2.0.
To generate candidate responses, we used LLaMA-13B Touvron et al. . We considered the closed-book setting, where the model does not have access to supporting text for answering the questions. We performed experiments in the few-shot setting by providing 32 example question-answer pairs sampled from the training set.
A truncated prompt used for generating answers on the TriviaQA dev set is reproduced as an illustration in Figure F.1. Please note that the actual prompt used in the experiment contains 32 question-answer pairs.
For generating answers in the open-domain question answering task, we use the default Transformers parameters reported in the previous section. We extract an answer by considering the text until the first line break, comma, or period is encountered. We then normalize the answers: this involves converting the generated answers to lowercase, removing articles, punctuation, and duplicate whitespace. Generated answers are then compared using the exact match metric: an answer is considered correct only if it matches the provided answer exactly.
F.3 News summarization
We use the CNN/DM dataset [Hermann et al., 2015, See et al., 2017] that includes news articles from CNN and the Daily Mail paired with their human written summaries, and is available at https://github.com/abisee/cnn-dailymail under MIT License. We use the standard train set for finetuning, the validation set for selecting the best checkpoint, and the test set for all reported conformal experiments.
We use a T5 1.1 XL model, which includes roughly 3B parameters, and was further pretrained for 100k steps with a multilayer objective Schuster et al. [2022b]. We finetune the model on the train set for 200k steps with a batch size of 128 using 64 TPUv4 chips for approximately 40 hours. We use the Adafactor [Shazeer and Stern, 2018] optimizer with a deacy rate of , initial learning rate of and 1k warm-up steps.
To generate candidate responses, we use Nucleus sampling [Holtzman et al., 2020] with top-p set to , temperature , and maximum output length set to tokens.
To get the response components we use a simple sentence spliter and treat each sentence as a component. As a classifier for evaluating the correctness of each component, we use an independent T5 XXL model trained on a mixture of NLI datasets Honovich et al. , Schuster et al. [2022a]. Specifically, we leverage the model used in the TRUE benchmark Honovich et al. and is available at https://huggingface.co/google/t5_xxl_true_nli_mixture. This model was trained on SNLI [Bowman et al., 2015], MNLI [Williams et al., 2018], FEVER [Thorne et al., 2018], SciTail [Khot et al., 2018], PAWS [Zhang et al., 2019], and VitaminC [Schuster et al., 2021a] to make a binary prediction of whether an hypothesis sentence is entailed by the given premise (in three-way datasets, the neutral class was merged with the negative class). We query the model with each component as the hypothesis, and the source summary as the premise, and measure the log-probability of predicting “entailment”.
F.4 Length-normalization
For all tasks, we apply length-normalization [Wu et al., 2016] to the model logits, i.e. we compute:
Appendix G Additional results
We describe another metric useful to characterize the effectiveness of the components identified by our component selection method.
In particular, we observe that component sets generated using scoring functions based on an auxiliary Classifier outperform uncertainty measures based solely on the span logits provided by the model.
Appendix H Qualitative results
We present qualitative results for radiology report generation and news summarization. In this section, we use the Sum method and consider . The choice of and is reported in Table H.7. We use 30% of the dev dataset (chosen uniformly at random) to determine as described in §4.3, and reserve the remaining 70% of the dataset for qualitative inspection. The corresponding values of and are reported in Table H.7. Notably, the method produces for the CNN/DM task, indicating that individual summaries are not rejected based on their quality but only for redundancy reasons.
In Figure H.1, an X-ray example is shown, depicting left basilar opacities while the rest of the X-ray appears normal. Table H.1 indicates that our method terminates the generation process after producing three samples. The third generation correctly identifies “apical scarring”; however, it mistakenly attributes it to the right lung instead of the left lung. This highlights a limitation of using CheXbert as the basis for the admission function, as its label granularity does not differentiate between left and right. Our component selection method accurately identifies several sentences that align with the reference report. These sentences are displayed in bold. Notably, our method avoids emphasizing low-confidence findings such as “right apical scarring” and instead focuses on the absence of an acute cardiopulmonary process.
A more challenging example is described in Figure H.2. The report mentions an enlarged heart, signs of cardiomegaly, and edema. Samples 4 and 5 correctly capture these findings but are considered incorrect due to the inclusion of “effusion.” The conformal selection of components chooses not to highlight any sentences since none of them meet the confidence threshold defined by .
In Tables H.3–H.6, we illustrate how our method continues sampling candidate summaries until the produced set is deemed acceptable. Specifically, Table H.3 demonstrates that the component selection process highlights the main idea while excluding minor ideas, which exist in multiple variations. Table H.4 exemplifies that the method stops after Sample 9, not because Sample 9 has the highest score, but because the sum of the scores collectively exceeds the target threshold of . Indeed, as shown in Table H.5, a higher individual score does not necessarily imply that a generation is more acceptable than one with a lower score. Finally, Table H.6 reveals a model failure, where the scores indicate high confidence in Sample 2, but the proposed generations are missing some main ideas from the reference summary.