Skip to content

Distilling What Matters: Confidence-Aware Selective Distillation for Large Language Models

Conference: NeurIPS2026 (task-package assignment; the cache is arXiv v1)
arXiv: 2609.36734
Area: Model Compression
Keywords: knowledge distillation, confidence gating, selective supervision, epistemic uncertainty, bidirectional KL divergence

TL;DR

CaRE-KD selects Forward or Reverse KL per token using relative teacher–student confidence and uses MC-dropout BALD to reject supervision when the teacher is uncertain but the student is comparatively certain, improving distillation quality in the tested settings and reaching 62.4 on MBPP and 73.9 on GSM8K, although gains are not universal and training becomes substantially more expensive.

Background & Motivation

Generative knowledge distillation usually treats a larger model's next-token distribution as the learning target for a smaller model. Forward KL encourages the student to cover the teacher's probability mass, transferring a rich set of candidate answers but also exposing the student to unreliable tails. MiniLLM uses Reverse KL, Distillm uses skewed divergences, and ABKD changes probability-mass allocation through a parameterized divergence. These approaches reshape imitation, but typically retain one geometric preference across all tokens.

A teacher's stronger average performance does not make it more reliable in every context. A student may already have a sharp distribution while the teacher spreads probability over contradictory continuations; compulsory matching can dilute the student's existing judgment. Conversely, when the student is uncertain and the teacher has a clear prediction, uniformly mode-seeking Reverse KL can lose useful coverage. The paper calls the first mismatch the fidelity trap, but a sharp distribution does not automatically imply a correct prediction. This distinction limits how the mechanism should be interpreted.

Rather than cleaning the entire dataset beforehand, the method answers two separate questions during training: which distillation direction to use for this token, and whether to accept this supervision at all. Core Idea: use relative entropy confidence to select token-level learning geometry, then use relative teacher–student epistemic uncertainty to decide whether to update, separating how to imitate from when to trust the teacher.

Method

Overall Architecture

CaRE-KD changes the training objective and update selection, not the student's network architecture. The teacher and student produce next-token distributions on identical prefixes. “Confidence-Gated Divergence” supplies token-level losses, while “Revival Epistemic Rejection” separately estimates sequence-level reliability through stochastic forward passes and decides whether to retain the supervision.

Under student-generated-output (SGO) training, the student generates continuations and the teacher scores those same on-policy token contexts. Non-SGO training uses teacher-forced contexts from ground-truth responses. The gate and distillation loss must share the same prefix; otherwise their confidence comparison concerns different conditional predictions and cannot identify which model is more trustworthy for the current token.

The diagram shows training dependencies, not mandatory deployment modules. BALD estimation and token losses can be computed separately; their connection indicates whether the loss is allowed to drive an update. Deployment uses only the trained student, without the teacher, stochastic uncertainty passes, or rejection rule.

%%{init: {'flowchart': {'rankSpacing': 24, 'nodeSpacing': 28, 'padding': 6, 'wrappingWidth': 400}}}%%
flowchart TD
    A["Training prefix<br/>Teacher and student distributions"] --> B["Confidence-Gated Divergence"]
    B --> C["Revival Epistemic Rejection"]
    A -. "MC dropout and historical quantiles" .-> C
    C -->|Retain supervision| D["Update student parameters"]
    C -->|Reject supervision| E["Skip the corresponding update"]
    D -. "Training complete" .-> F["Inference: student-only generation"]

Key Designs

1. Confidence-Gated Divergence: select the imitation direction using local relative confidence

Single-pass confidence is defined by normalized entropy: concentrated probabilities imply higher confidence. It requires neither an auxiliary confidence predictor nor a model's verbal report of certainty. When teacher confidence is at least student confidence, hard gating selects Forward KL to cover the teacher distribution. When the teacher is less confident, it selects Reverse KL to reduce imitation of the teacher's tail. The comparison is relative: even an absolutely confident teacher can enter the reverse branch if the student is more confident.

Let \(p\) be the teacher distribution, \(q\) the student distribution, \(H\) Shannon entropy, \(\mathcal{V}\) the vocabulary, and \(m\) the soft-gate margin offset. The mechanism in main-text Equations (1)–(3) is:

\[ c(p)=1-\frac{H(p)}{\log|\mathcal{V}|},\qquad \mathcal{L}_{\mathrm{CARE}}=g\,\mathrm{KL}(p\|q)+(1-g)\,\mathrm{KL}(q\|p), \qquad g_{\mathrm{hard}}=\mathbb{I}[c(p)\geq c(q)],\qquad g_{\mathrm{soft}}=\sigma(c(p)-c(q)-m). \]

Hard gating actually selects one direction; soft gating continuously weights both, so it should not be described as assigning every token completely to a single branch. The margin shifts the sigmoid's center without changing its slope, and a small margin therefore does not turn soft gating into hard gating. Without an additional slope parameter, a finite confidence gap need not bring the sigmoid close to either endpoint.

Forward KL penalizes omission of teacher probability mass, whereas Reverse KL pays more attention to locations where the student currently places probability, reducing the influence of teacher tails. Nevertheless, Reverse KL still moves toward the teacher distribution: it does not preserve the original student distribution and cannot guarantee that the student's confident prior is correct. The second design supplies the actual mechanism for stopping supervision.

2. Revival Epistemic Rejection: stop supervision when stochastic teacher predictions are unstable but the student is stable

High single-pass entropy may simply reflect multiple valid answers; it does not independently identify parameter-level uncertainty. Revival runs MC dropout for the teacher and student separately, computes token-level BALD, and averages across tokens to obtain a sequence-level score. BALD subtracts mean individual predictive entropy from the entropy of the mean prediction. If each pass is sharp but different passes select different answers, the first term remains high and the second low, exposing disagreement among stochastic model realizations.

Let \(\omega\) denote a stochastic dropout realization, \(\mathcal{I}_{T}\) and \(\mathcal{I}_{S}\) the teacher and student sequence-level BALD scores, and \(Q_{\tau}\) each model's running quantile of historical scores. Main-text Equations (4)–(5) are:

\[ \mathcal{I}_{\mathrm{BALD}}(\mathbf{x})= H\!\left[\mathbb{E}_{\omega}p(y\mid\mathbf{x},\omega)\right] -\mathbb{E}_{\omega}\!\left[H(p(y\mid\mathbf{x},\omega))\right], \qquad M_{\mathrm{Revival}}=\mathbb{I}\!\left[ \mathcal{I}_{T}>Q_{\tau}(\mathcal{I}_{T})\ \land\ \mathcal{I}_{S}<Q_{\tau}(\mathcal{I}_{S})\right]. \]

The rule does not directly test whether teacher BALD exceeds student BALD. Instead, each model is compared with its own history. Rejection requires the teacher to lie in its relatively uncertain region and the student in its relatively certain region; uncertainty in both models does not automatically cause rejection. This avoids rejecting all supervision merely because an example is difficult and reduces sensitivity to different absolute BALD scales.

Appendix §8.1 explicitly states that uncertainty estimation forces dropout layers into training mode and sets attention dropout to 0.1 even when the original deployment configuration disables dropout, restoring the configuration afterward. The default is 3 stochastic forward passes. This reliability estimate therefore has additional computational cost rather than being a free byproduct of single-pass entropy gating; teacher parameters remain fixed during distillation.

The granularity requires care. Section §2.2 and Appendix §8.1 define sequence scores, but main-text Equation (6) places the rejection mask outside the entire batch sum, and the discussion describes skipping batch back-propagation. The cache does not specify how multiple sequence masks become a batch decision. Its intention to reject corresponding supervision and prevent related updates is clear, but the note cannot replace Equation (6) with a particular per-example weighted average and present that reconstruction as the author's implementation.

The paper progressively strengthens rejection so that training begins with broad supervision and becomes more selective later. It uses \(\tau\) both as a running quantile and as a “target skip percentage,” but these quantities are not generally identical: the joint condition's activation rate depends on both score distributions and their correlation. The reported favorable 50–70% range should be read as an experimental scheduling setting, not an exact rejection rate derived directly from Equation (5).

Loss & Training

When supervision is retained, the student optimizes the gated distillation objective; when rejected, the corresponding update is skipped. Hard gating is a detached step function with zero gate gradient, retaining only per-token branch selection. A soft gate that remains connected to the computation graph introduces a gate-selection term in addition to the two KL learning gradients. Proposition 3.1 and Appendix §9.1 give:

\[ \nabla_{\theta}\mathcal{L}_{\mathrm{CARE}} =g\nabla_{\theta}\mathcal{L}_{\mathrm{FKL}} +(1-g)\nabla_{\theta}\mathcal{L}_{\mathrm{RKL}} +(\mathcal{L}_{\mathrm{FKL}}-\mathcal{L}_{\mathrm{RKL}})\nabla_{\theta}g, \qquad \nabla_{\theta}g=-g(1-g)\nabla_{\theta}c(q_{\theta}). \]

Here \(\theta\) denotes student parameters. When Forward KL exceeds Reverse KL, the extra term can increase student confidence under gradient descent. This should not be attributed to hard gating as a confidence-gradient effect. Experimentally, hard gating achieves the highest peak performance, while soft gating offers greater stability across configurations. The theoretical explanation and the strongest experimental configuration therefore need separate descriptions.

Proposition 3.2 distinguishes token-dependent branch selection from a single global skew parameter, which generally cannot reproduce that position-dependent behavior. It does not prove that static divergences lose on every task. Theorem 3.5 has explicit conditions: optimization stays in a fixed branch and the student converges to the teacher, after which student entropy converges to teacher entropy. It does not guarantee global convergence while branches switch, monotonic entropy changes, or calibration against ground-truth answers.

The entropy-as-BALD theorem also requires two assumptions: each stochastic predictive distribution is sharp, and the single-pass distribution approximates the mean stochastic prediction. Multiple valid continuations, high aleatoric uncertainty, or substantial stochastic predictive disagreement can weaken these approximations. Every reported result uses full MC-dropout BALD; inexpensive single-pass entropy rejection is not an experimentally validated equivalent replacement.

Appendix §10 uses LoRA rank 16 and learning rate \(5\times10^{-5}\). On Dolly, students below 1B parameters use batch size 32 for up to 20 epochs; students above 1B use batch size 8 for 10 epochs, selecting checkpoints by validation ROUGE-L. UltraChat, WizardCoder, and MetaMathQA use 3, 2, and 2 epochs respectively. BALD defaults to 3 stochastic passes and temperature 3.0. Instruction evaluation uses temperature 0.8, top-p 0.95, and at most 512 tokens; code and math use greedy decoding with at most 1024 tokens. Experiments run on one A100 80GB GPU.

Key Experimental Results

Main Results

Evaluation covers 8 teacher–student pairs and 11 benchmarks. The following table extracts average SGO instruction results from main-text Table 1. Cells contain “ROUGE-L / LLM-judge factuality,” with GPT-5-Mini assigning a 0–100 reference-consistency score. Results are averaged over 5 seeds as reported. SRKL means Skewed Reverse KL.

Teacher → Student FKL SRKL CaRE-Div
GPT2-XL → GPT2-base 17.0 / 16.7 19.3 / 19.5 19.6 / 21.0
GPT2-XL → GPT2-large 16.1 / 15.7 17.9 / 16.9 18.1 / 19.1
OPT-2.7B → OPT-125M 16.1 / 15.5 19.0 / 20.2 19.3 / 18.5
Gemma-2-9B-IT → Gemma-2-2B-IT 16.3 / 18.2 16.5 / 20.6 17.1 / 21.6
OpenLLaMA2-7B → OpenLLaMA2-3B 22.3 / 27.9 26.4 / 35.3 26.2 / 33.8

OPT has higher ROUGE-L but lower factuality than SRKL; OpenLLaMA does not exceed SRKL on either average metric. Table 1's caption says average ROUGE-L is best on 3 pairs, whereas subsequent prose says 4. The row values support the latter. This note preserves the values and flags the textual conflict rather than claiming universal superiority.

Specialized-domain results from main-text Table 3 follow. Code uses execution-based pass@1; math uses exact-match accuracy according to §4 and §10.2. Table 3's caption labels all cells ROUGE-L / LLM, while its math column headers say pass@1, both inconsistent with the metric description.

Dataset Metric Original student Distillm2 ABKD GKD Distillm CaRE-KD
HumanEval pass@1 32.3 43.3 42.7 43.1 42.9 43.3
MBPP pass@1 58.5 60.3 59.8 60.2 60.0 62.4
GSM8K exact-match accuracy 69.9 72.2 71.0 71.4 71.7 73.9
CollegeMath exact-match accuracy 37.1 44.2 44.3 43.9 44.1 46.1

Relative to the strongest listed baseline, MBPP improves by 2.1 percentage points, GSM8K by 1.7, and CollegeMath by 1.8. HumanEval only ties the best result. The paper's prose is not fully consistent about the CollegeMath comparator, so these gains are calculated from Table 3's values rather than mixing baseline labels.

Ablation Study

Main-text Table 2 isolates Revival through the change in ROUGE-L, defined as “with Revival − without Revival.” The table below retains averages over the four instruction benchmarks and necessary counterexamples, without treating ablation deltas as absolute scores.

Student Loss Average change Per-task counterexample / note
GPT2-base CaRE +0.90 Positive on all four tasks
GPT2-large CaRE +0.46 Dolly −0.04
OPT-125M CaRE +1.42 Super-Natural Instructions +3.62
Gemma-2-2B-IT CaRE +2.90 Vicuna −0.40
OpenLLaMA2-3B CaRE +0.73 Positive on all four tasks
Gemma-2-2B-IT FKL −3.91 Filtering is not a universal gain
GPT2-large RKL −0.28 Self-Instruct −1.78

CaRE's average deltas are positive, but two individual cells are negative; static losses combined with Revival can even decline on average. The paper reports a one-sample test statistic of 3.98 with significance below 0.001 for CaRE's deltas. This supports average complementarity in the tested configurations, not a claim that filtering benefits every loss. Numeric curves from Figures 3 and 8 are absent from the text cache, so only the prose trends for gating, scheduling, and sample count are used, without invented pointwise ablation scores.

Key Findings

  • Gains are clearer at high compression: average ROUGE-L increases over FKL are 2.6 for GPT2-base and 3.2 for OPT-125M, but only 0.3 for each against the stronger SRKL baseline.
  • Factuality evidence concerns end-to-end outputs, not proven hallucination identification by BALD. Section §4.2 reports teacher factuality of 25.2 in the rejection region versus 28.6 outside it, but Wilcoxon significance is 0.06 and another region comparison yields 0.26; neither is significant.
  • Training overhead is not universally modest. At the default 3 passes, §5 reports approximately 42% extra time for instruction following, 106% for chat, 170% for code, and 148% for math. A Dolly epoch increases from 5520 to 7860 seconds. Deployment gains no additional inference modules, but offline training requires substantially more compute.

Highlights & Insights

  • Relative confidence is finer-grained than assuming the teacher is always right. The transferable idea is not simply rewarding sharp distributions: it is adapting supervision geometry to the current context while retaining an explicit rejection option.
  • Gating and rejection have different responsibilities. Gating continues to learn from the teacher while changing coverage versus mode selection; rejection stops particular supervision from influencing the student. This distinction gives their combination more explanatory value than another global divergence parameter.
  • Historical quantiles do not require teacher and student BALD to share an absolute scale. Reuse should nevertheless record the quantile window, actual activation rate, and retained-example distribution rather than publishing only a target skip setting.

Limitations & Future Work

  • Confidence is not accuracy, and low BALD can accompany consistently wrong predictions. Independent correctness checks or stronger calibration experiments are needed to determine whether the protected student priors are actually correct.
  • Appendix Table 9 gives default CaRE-KD ECE as 0.34 versus FKL's 0.22; the softened configuration has 0.28, not the “same ECE” claimed in the prose. Single-reference token matching has limitations, but the reported values do not support best ground-truth calibration.
  • Sequence-level definitions conflict with batch-level Equation (6), and the mapping from quantile parameter to target skip percentage is insufficiently explicit. Reproduction should verify mask aggregation, padding handling, and the actual rejection schedule rather than filling these gaps from this note.
  • Nonzero disagreement after forcibly enabling MC dropout does not establish calibrated epistemic uncertainty. Appendix Table 4 checks only GPT2-XL and OPT-2.7B and cannot be generalized to every teacher model.
  • Replacing BALD with single-pass entropy has conditional theoretical justification only; all experiments still use full BALD. Equal-compute comparisons among the fallback, extra training steps, and full CaRE-KD are needed to establish whether the additional forward passes are worthwhile.
  • Nonsignificant hardness and diversity differences in Appendix Table 10 mean only that the OpenLLaMA-Dolly check detected no difference, not that all tasks are free of selection bias. Chat comparisons with GKD and Distillm also remain missing.
  • vs MiniLLM: MiniLLM suppresses teacher tails using Reverse KL; CaRE-KD selects Forward or Reverse geometry per token and rejects updates when reliability is inadequate. Its cost is more complicated control logic and stochastic forward computation.
  • vs Distillm / ABKD: Skewed KL and parameterized divergences use global geometry parameters; this method uses context-dependent gates. OpenLLaMA's average result remains slightly below Distillm, showing that dynamic selection is not an unconditional replacement for static methods.
  • vs GKD: On-policy data addresses mismatch between training contexts and student generation, while this method addresses teacher trust within a given context. The approaches can be combined; in SGO training the gate and teacher scoring must share the student's same generation trajectory.
  • Research direction: Test rejection's side effects where the student is confident but wrong, and add code unit tests or math answer verification as trust signals. Match baselines by compute budget rather than epoch count. These are experimental suggestions from this note, not findings already validated by the paper.

Rating

  • Novelty: 4/5; token-level geometric selection and sequence-level reliability rejection form a clear combination.
  • Experimental Thoroughness: 3/5; broad model and task coverage, but equal-budget comparisons, factuality validation of rejection, and some baselines are missing.
  • Writing Quality: 3/5; the mechanism is well explained, but mask granularity, metric labels, and several interpretations of numerical results conflict.
  • Value: 4/5; useful for offline high-compression distillation, provided teacher reliability and additional training cost are assessed together.