Coupled Calibration and Learning: Mitigating Teacher Bias in
arXiv:2609.17474v1 [cs.LG] 15 Sep 2026
LLM Distillation without Target-Domain Reward Feedback Haichen Hu♮
Yuheng Zhang†
David Simchi-Levi∗
MIT
UIUC
MIT
Abstract Large language model (LLM) distillation aims to transfer the capabilities of a powerful teacher to a smaller student. Direct imitation, however, can also transfer the teacher’s systematic bias and errors. This challenge is particularly pronounced under covariate shift, when the teacher’s reliability on target questions is uncertain and targetdomain reward feedback is unavailable. We propose Coupled Calibration and Learning (CCL), an LLM distillation algorithm that couples teacher calibration with student updates through token-level branching, using reward feedback only on source questions. Each iteration calibrates the teacher using source feedback and then uses the calibrated teacher to train the student on target questions. The updated student, in turn, informs subsequent calibration. In an autoregressive policy framework, we prove that the output student’s expected average Kullback-Leibler divergence to the oracle student converges to zero at a polynomial rate in the number of iterations. The oracle maximizes the true reference-regularized target reward within the student class, which need not represent the unrestricted optimal policy. Our analysis quantifies the progress of projected student gradient updates while controlling the error in teacher calibration. We further establish a separation from regularized direct matching: its error relative to the oracle student can remain bounded away from zero even when the teacher achieves higher regularized target reward than every student policy. These results demonstrate that LLM distillation can overcome persistent teacher bias and recover the optimal student through coupled calibration and learning, without target-domain reward feedback.
1
Introduction
Large language models (LLMs) have become a central component of modern artificial intelligence, with applications ranging from language understanding and content generation to mathematical reasoning and program synthesis (OpenAI, 2023; Gemini Team, 2023; Grattafiori et al., 2024; Qwen Team, 2024; DeepSeek-AI, 2025). As these models increasingly supply predictions, decisions, and training data for other systems, understanding their statistical behavior has become an important research problem. ♮
Email: [email protected] Email: [email protected] ∗ Email: [email protected] †
1
Statistical learning theory provides a principled framework for understanding the capabilities and limitations of modern AI. It studies when overparameterized predictors generalize, how model predictions can support valid statistical inference, and how to evaluate black-box prediction procedures. For language models, theoretical analyzes also examine how pretraining benefits downstream tasks, how transformers learn from examples in context, and how preference feedback guides policy learning. These questions have motivated progress in generalization theory, statistical inference, and the analysis of language-model training (Bartlett et al., 2020; Angelopoulos et al., 2023; Saunshi et al., 2021; Bai et al., 2023; Kim et al., 2024a; Ye et al., 2024; Wainwright, 2025; Hu and Simchi-Levi, 2025b,a, 2026). Within this broader program, a central challenge is to explain when information from an existing model can support the reliable training of another model. Knowledge distillation is a prominent approach to this challenge. A larger or more capable teacher supplies output probabilities or generated responses to train a student, often with substantially lower deployment costs. This idea underlies both classical model compression and recent methods for transferring language and reasoning capabilities to smaller models (Hinton et al., 2015; Gu et al., 2024; DeepSeek-AI, 2025). The teacher provides a rich source of synthetic supervision, making distillation particularly attractive when collecting task-specific demonstrations is difficult. A central difficulty, however, is the propagation of teacher-bias. Matching the teacher’s predictions can reproduce its systematic errors as well as its useful knowledge. For example, image-classification experiments show that distillation can amplify teacher errors on difficult classes even when average student accuracy improves (Lukasik et al., 2022). This concern is especially relevant under covariate shift, where the target questions differ from the source questions on which reliable supervision is available. Strong source performance does not by itself certify the teacher’s accuracy on the target questions. Related experiments with LLMs show substantial accuracy losses when in-context demonstrations and evaluation examples come from different topic domains (Roussinov et al., 2025). Nevertheless, the teacher may retain useful knowledge acquired from source data. The statistical challenge is therefore to exploit that information while correcting the errors that direct imitation can propagate (Yamamoto and Wainwright, 2026). The absence of reliable target feedback makes this problem more difficult. A reward model or verifier developed for one collection of questions need not remain reliable on another: multilingual evaluations, for example, find lower reward-model accuracy and inconsistent preferences across languages (Gureja et al., 2025). Constructing new feedback can also require substantial resources. Training mathematical process verifiers has involved extensive human step annotations (Lightman et al., 2024), while executable code evaluation relies on task-specific tests and execution infrastructure (Chen et al., 2021). These considerations motivate a setting in which trusted reward feedback is available on source questions, but only the questions themselves are available in the target dataset. The learner must then use source feedback to address target teacher bias, without evaluating target answers against their true rewards. This leads to the following question: Can we distill a near optimal student from a biased teacher under covariate shift, using reward feedback only on source questions and none on the target questions? Our contribution. We answer this question affirmatively by proposing Coupled Calibration and Learning (CCL), an LLM distillation algorithm that learns from a biased teacher using reward feedback only on source questions. In an autoregressive policy framework, our benchmark is the oracle student: the policy within the student class that maximizes the true target reward with KL regularization to a pre-trained reference policy. 2
Our algorithm couples teacher calibration with student updates in an iterative procedure. Each iteration calibrates the teacher using source reward feedback and then uses the calibrated teacher to train the student on target questions. The updated student, in turn, informs the next calibration step so that teacher calibration and student learning proceed together throughout training. We prove that the expected average KL divergence from the student returned by CCL to the oracle student converges to zero at a polynomial rate in the number of iterations. The analysis uses student gradient progress near the oracle and exploration outside a fixed near-optimal region, while controlling calibration and finite-rollout errors. We further establish a separation: direct distillation can retain a strictly positive KL divergence between the oracle student and the trained student policy even when the teacher achieves a higher true regularized reward than every policy in the student class. Thus, our method can eliminate a persistent error of direct imitation and recover the optimal student without target-domain reward feedback. Paper structure. Our paper is organized with the following structure. Section 3 introduces the distillation problem, the autoregressive policy model, and the oracle student benchmark. Section 4 presents our algorithm and explains how it couples teacher calibration with student updates. Then, Section 5 establishes convergence to the optimal oracle student and outlines the main steps of the analysis. Section 6 establishes a separation from regularized direct teacher matching to show that our bound is strictly better than direct distillation. Notation. We write EX [·] and PX (·) for expectation and probability with respect to the random variable X, respectively. Conditional expectation and probability are denoted by E[· | F ] and P(· | F ) for a sigma-algebra F. We use σ (X1 , . . . , Xk ) for the sigma-algebra generated by (Xi )ki=1 , and F ∨ G for the smallest sigma-algebra containing both F and G. For a vector v, ∥v∥2 denotes its Euclidean norm; for a matrix A, ∥A∥op denotes its induced Euclidean operator norm. For probability distributions P , Q on a common finite set Z, their Kullback-Leibler divergence P is denoted by KL(P ∥Q) = z∈Z P (z ) log(P (z )/Q(z )), where terms with P (z ) = 0 are zero, and the divergence is +∞ if P (z ) > 0 = Q(z ) for some z. For a nonempty closed convex set C, ProjC (v ) := arg minu∈C ∥u − v∥22 denotes Euclidean projection. The notation Unif (C ) denotes the uniform distribution over C. For two sequences (ak )k≥1 and (bk )k≥1 with bk > 0, we write ak = O (bk ) if there exist constants C > 0 and k0 such that |ak | ≤ Cbk for all k ≥ k0 , and ak = o(bk ) if |ak |/bk → 0 as k → ∞.
2
Related Work
Statistical theory of distillation and imitation learning. Statistical analyzes of distillation study the benefits of teacher-generated supervision and the propagation of teacher error. Menon et al. (2021) explain the bias–variance tradeoff of soft labels, while Ildiz et al. (2025) characterize high-dimensional distillation risk under model and covariate shift. Xie et al. (2026) identify settings where a risk-minimizing teacher preserves the restricted student’s optimum and improves the statistical efficiency of averaged SGD. To correct imperfect supervision, Dao et al. (2021) develop cross-fitting and loss corrections, and Iliopoulos et al. (2022) analyze student-dependent reweighting of noisy pseudo-labels. Relatedly, Xia and Wainwright (2024) construct pseudo-responses using training-only helper covariates and obtain prediction bounds combining an oracle rate with surrogate error. For sequential prediction, Ross et al. (2011) establish a no-regret foundation
3
for learning from expert feedback at learner-visited states, while Czarnecki et al. (2019) analyze policy-distillation updates and convergence in tabular settings. Foster et al. (2024) show that, under realizability, online expert access need not improve worst-case statistical complexity over offline behavior cloning with logarithmic loss. With noisy expert feedback, Sriraman et al. (2026) establish an offline on-policy separation for learning a realizable clean expert. Beyond realizability, Zhang et al. (2026) study how student misspecification and alignment between expert scores and rewards affect the benefits of online imitation, and give finite-sample guarantees using base-policy sampling. Closest to our motivation, Yamamoto and Wainwright (2026) study bias propagation under source-target covariate shift. Their method refits teachers to student residuals and achieves a provable separation from direct soft matching. Our work instead couples source-reward-based teacher calibration with target-side LLM distillation. We establish convergence to the regularized oracle within the student class and separation from regularized direct matching, without target reward feedback or requiring the student to represent the unrestricted optimal policy. Reinforcement learning and LLM post-training. Reinforcement learning plays an important role in modern deep learning, with a substantial theoretical literature on exploration, policy optimization, and learning with function approximation (Jiang et al., 2017; Jin et al., 2020; Agarwal et al., 2021; Xie et al., 2021, 2023; Zhang et al., 2023; Qian et al., 2024; Hu et al., 2026). For preference-based learning, Zhu et al. (2023) connect reward estimation to policy performance and establish guarantees for pessimistic learning. Xiong et al. (2024) develop offline, online, and hybrid algorithms for KL-regularized preference learning with finite-sample guarantees. Xie et al. (2025) introduce exploration bonuses for provably sample-efficient preference optimization, while Zhao et al. (2025) characterize how KL regularization and reference-policy coverage affect sample complexity. For LLM post-training, Chen et al. (2026) study how pre-training provides response coverage for downstream improvement, and Foster et al. (2025) distinguish the statistical and computational roles of base-model coverage. Huang et al. (2025) analyze self-improvement through sharpening, where training amortizes the selection of high-likelihood responses. Related inference-time analyses characterize Best-of-N through win-rate guarantees (Sriraman and Block, 2026) and use pessimistic scoring to mitigate reward hacking (Yu et al., 2026). For outcome-supervised learning, Jia et al. (2025) connect outcome feedback to process-level learning, Yuan et al. (2025) develop trajectory Bellman residual minimization for KL-regularized policy learning, and Chen et al. (2025) establish sample-complexity guarantees for outcome-based online RL. Kim et al. (2026) analyze coverage improvement and convergence in on-policy preference learning and reward distillation, explicitly accounting for reward-model error in the latter. These works study how feedback and coverage support policy improvement and response selection. Our work addresses biased-teacher LLM distillation under source-target covariate shift. We couple source-reward-based teacher calibration with target-side student learning and prove convergence in average target KL to the regularized oracle within the student class, together with separation from regularized direct matching, without target reward feedback. Transfer learning under covariate shift. Transfer-learning theory studies how source supervision supports prediction on a different target distribution. Under covariate shift, linear and kernel regression analyzes quantify this transfer through source-target covariance geometry and distributional overlap (Lei et al., 2021; Ma et al., 2023). In well-specified parametric models, Ge et al. (2024) establish minimax guarantees for source-only maximum likelihood estimation, with transfer difficulty governed by source and target Fisher information rather than a bounded density ratio. Pseudo-labeling methods further use unlabeled target covariates to select source-trained estimators 4
in kernel ridge regression and kernel generalized linear models (Wang, 2026; Weill and Wang, 2026). Beyond covariate shift, Xia and Klusowski (2026) study oversampling under label shift, separating balanced-data risk from the cost of estimating the minority-class distribution. These results characterize prediction under distribution shift, but do not study LLM distillation without target feedback. Under shared policy realizability and joint source identification, our analysis controls the student’s error on fixed target questions using reward feedback confined to source questions.
3
Model setup
In this section, we formulate LLM distillation for post-training under a fixed source–target design. We model autoregressive generation as a finite-horizon, token-level Markov decision process M = (S, A, P , R, H ), where S is the prefix-state space, A is a finite token vocabulary, P is the deterministic transition kernel, R is the terminal answer reward, and H ≥ 1 is the generation horizon. Denote X as a context space representing the set of potential questions. Each x ∈ X represents a question presented to the LLM, including its instructions and any accompanying context; we refer to this complete input as a prompt. Generation starts at s1 = (x, ∅). At step h, the state sh = (x, a1:h−1 ) records the question and all previously generated answer tokens. The LLM acts as a policy π: it selects a legal token ah ∼ π (· | sh ) and appends it to the prefix, so P (s′ | sh , ah ) = 1{s′ = (x, a1:h )}, h = 1, . . . , H. The reward is evaluated on the completed answer and is specified below. We first make precise which tokens and answers are feasible. Definition 3.1 (Legal tokens and feasible trajectories). Let A contain EOS and null. For 1 ≤ h ≤ H, define the legal-token set at sh = (x, a1:h−1 ) by (
B ( sh ) : =
A \ {null}, EOS ∈ / {a1 , . . . , ah−1 }, {null}, EOS ∈ {a1 , . . . , ah−1 }.
The state space S consists of the prefixes generated from (x, ∅), x ∈ X , by repeatedly appending legal tokens, up to length H. States sH +1 are terminal, with B (sH +1 ) := ∅. For a reachable state sh with h ≤ H, define n
o
A(sh ) := bh:H ∈ AH−h+1 : bk ∈ B ((x, a1:h−1 , bh:k−1 )) for k = h, . . . , H . Here (x, a1:h−1 , bh:k−1 ) appends bh:k−1 to the existing prefix, with empty blocks omitted. Set A(sH +1 ) := {∅}, A(x) := A((x, ∅)). Thus B (sh ) contains individual legal tokens, A(sh ) contains feasible remaining token sequences, and A(x) contains full feasible answers. Generation ends at EOS or horizon H, with null padding after EOS. This convention gives every answer length H, and feasibility is independent of reward. A policy assigns probability zero to
5
illegal tokens; on each nonterminal state, its probabilities over B (sh ) sum to one. The deterministic transitions and successive token choices induce the full-answer law π (a1:H | x) =
H Y
π (ah | x, a1:h−1 ), a1:H ∈ A(x).
(3.1)
h=1
The same notation π therefore describes both the token policy and its induced distribution over complete answers. Post-training seeks to improve answer quality while retaining the behavior of a pretrained model. We study this task under covariate shift, with a fixed source dataset of training questions and a fixed target dataset of questions that the student is intended to answer: e j }m Dsrc = {xi }ni=1 , Dtar = {x j =1 , n, m ≥ 1. ej is the jth target question, each including its associated Here xi is the ith source question and x context. Reward supervision is available for candidate answers to source questions, whereas the learning objective concerns the student’s answers to target questions. The two datasets may contain different mixtures of question types. We condition throughout on these fixed datasets; randomness comes from policy sampling and the training algorithm.
The terminal reward is a deterministic function R : {(x, a1:H ) : x ∈ X , a1:H ∈ A(x)} −→ [0, 1]. We only have access to an exact verifier on the source prompts: given xi and any a1:H ∈ A(xi ), it returns R(xi , a1:H ). Source supervision is therefore supplied by evaluations of candidate answers, and the dataset itself contains only prompts. Reward feedback is unavailable on the target prompts; ej , a1:H ) denotes their latent true answer quality and is never queried during training. This reR (x striction reflects the difficulty of extending reliable reward evaluation to new questions. Developing a target-domain reward model can require expert-designed rubrics, labeled answers, and careful validation of the grading criteria. These requirements make target reward very costly to obtain. Let πpre be the frozen pretrained reference policy, with positive probability on every legal token. For λ > 0, our target post-training objective is
H m X ej , a1:h−1 ) 1 X π (ah | x . ej , a1:H ) − λ Jλ,m (π ) := Ea1:H ∼π (·|exj ) R(x log ej , a1:h−1 ) m j =1 πpre (ah | x h=1
(3.2)
The first term measures answer quality; the second penalizes deviation from the pretrained policy. By (3.1), the expected log-ratio sum equals the full-answer KL: " H X
#
π (ah | x, a1:h−1 ) KL(π (· | x) ∥ πpre (· | x)) = Ea1:H ∼π (·|x) log . πpre (ah | x, a1:h−1 ) h=1 Thus λ controls the tradeoff between reward and proximity to the reference. The objective specifies the desired target behavior, although its reward term is unavailable during training. We address this information constraint through teacher distillation: source reward feedback calibrates an available teacher model, whose likelihoods then supervise the student on the target prompts. To formalize the teacher, its calibration, and the student, we use linear-softmax policy classes. Fix
6
known feature maps ϕ : S × A → RD and ϕstu : S × A → Rd , where D, d ≥ 1, with ∥ϕ(s, a)∥2 ≤ 1, ∥ϕstu (s, a)∥2 ≤ 1. The features may depend on the entire prefix, preserving the autoregressive dependence of the policy. For every nonterminal state s and legal token a ∈ B (s), define exp(θ⊤ ϕstu (s, a)) exp(w⊤ ϕ(s, a)) P , π ( a | s ) = . stu,θ ⊤ ⊤ b∈B (s) exp(w ϕ(s, b)) b∈B (s) exp(θ ϕstu (s, b))
πw ( a | s ) = P
(3.3)
Both policies assign zero probability to illegal tokens. For B > 0, the calibration parameter belongs to a nonempty compact convex set W , and the student parameter belongs to the closed ball Θ: W ⊆ {w ∈ RD : ∥w∥2 ≤ B}, Θ := {θ ∈ Rd : ∥θ∥2 ≤ B}. The given teacher is πtea = πwtea with wtea ∈ W . Its proposal policy remains frozen, while calibration adjusts a separate parameter within W , initialized at wtea . The teacher may be biased relative to the optimal target behavior. The reference πpre need not belong to either softmax class, and we impose neither sparsity nor an upper bound relating D to H. For the performance metric, our benchmark is the best post-trained policy within the student class. Specifically, we choose † θλ,m ∈ argmax Jλ,m (πstu,θ ). θ∈Θ
† θλ,m exists because Θ is compact, and the objective is continuous in θ. We call πstu,θ†
an oracle
λ,m
student: it optimizes the true target objective using the same class available to the learned student. For comparison, at each source or target prompt x, we define the unrestricted optimal policy by
πλ⋆ (· | x) ∈
H X
π (ah | x, a1:h−1 ) argmax Ea1:H ∼π (·|x) R(x, a1:H ) − λ log . π pre (ah | x, a1:h−1 ) π (·|x)∈∆(A(x)) h=1
(3.4)
Here ∆(A(x)) is the simplex of distributions over full feasible answers. Token conditionals are obtained from prefix marginals. Lemma B.1 establishes that this optimum is unique and assigns positive probability to every feasible answer. The student class may be unable to represent this unrestricted optimum. We connect source reward information to target behavior through the following realizability assumption. Assumption 3.2 (Realizability). There exists a single parameter wλ⋆ ∈ W such that, at every legal nonterminal prefix state s of every source or target prompt, πλ⋆ (a | s) = πwλ⋆ (a | s) for every a ∈ B (s). This assumption provides the shared structure needed for transfer: source and target prompts use the same optimal calibration coefficients, evaluated through their respective prefix features. Realizability is imposed on the calibration class and allows the student class to remain misspecified. Our fixed-design analysis uses this shared structure rather than a density-ratio assumption.
7
For a calibrated parameter w ∈ W , define the target distillation cost m λ X ej ) ∥ πw (· | x ej )) . KL(πstu,θ (· | x m j =1
Cw ( θ ) : =
(3.5)
Its integrand is computable from the student and calibrated-policy likelihoods, so the cost can be estimated using student rollouts on the fixed target prompts. Under realizability, Lemma B.1 gives argmax Jλ,m (πstu,θ ) = argmin Cwλ⋆ (θ ). θ∈Θ
θ∈Θ
This identity makes the role of calibration explicit: at the true calibration parameter, distillation targets exactly the regularized oracle student. We impose the following uniqueness assumption on the oracle student. † Assumption 3.3. The function Cwλ⋆ has a unique minimizer over Θ, argminθ∈Θ Cwλ⋆ (θ ) = {θλ,m }. † This assumption uniquely determines the oracle parameter θλ,m and hence its induced answer law on every target question. Neither interiority nor a positive student Hessian is required. For a learned calibration wt and student θt , write Ct (θ ) := Cwt (θ ) and define
εt := Ct (θt ) − min Ct (θ ), θ∈Θ
∆λ,m (θ ) := C
⋆ wλ
† (θ ) − Cwλ⋆ (θλ,m ),
(3.6)
m 1 X ej ) πstu,θ† (· | x ej ) . KL πstu,θ (· | x Kλ,m (θ ) := λ,m m j =1
The first quantity measures optimization error for the current calibrated objective; the second measures excess cost under the true calibration; and the third measures the average KL from the learned student to the oracle student. Our goal is to drive E[Kλ,m (θT )] to zero. This compares policies within the student class, allowing the minimum distillation cost itself to remain positive. Finally, source comparisons must contain enough information to identify the shared calibration parameter. For a teacher-generated prefix si,h = (xi , ai,1:h−1 ), compare its next token ai,h with a legal alternative b. The difference ϕ(si,h , ai,h ) − ϕ(si,h , b) determines the calibration direction observed in that comparison. For each h = 1, . . . , H, define the source information matrix Gh by Gh : =
n 1X
n i=1
Eai,1:h ∼πtea (·|xi )
1 |B (si,h )|
X
ϕ(si,h , ai,h ) − ϕ(si,h , b) × ϕ(si,h , ai,h ) − ϕ(si,h , b)
⊤
b∈B (si,h )
(3.7) The expectation is over the teacher’s output ai,1:h , and the inner average is uniform over legal tokens. We impose the following identification condition on their average across token positions. Assumption 3.4. Define the joint source information matrix as Gjoint := H1 that it is positive definite, i.e., µjoint := λmin (Gjoint ) > 0.
PH
h=1 Gh . We assume
This assumption requires comparisons across all token positions to identify every calibration direction; individual matrices Gh may be singular. At a padded state, the only legal token is null, so its feature difference and information contribution are zero. 8
,
4
Algorithm
In this section, we present the CCL algorithm and provide a detailed explanation of its steps. The algorithm couples two components: teacher calibration using source reward feedback, and student distillation using the calibrated teacher on target prompts. Calibration adjusts the teacher’s predictions toward the regularized optimal policy, while distillation uses these adjusted predictions to train the student. The student also participates in calibration by supplying alternatives to the teacher’s proposed tokens. These components therefore interact throughout training: calibration changes the student’s training objective, and the updated student changes the comparisons used for subsequent calibration. The central mechanism for estimating the calibration update is token-level branching. We select a token position in a teacher-generated source answer and retain the preceding prefix. At that prefix, we pair the teacher’s next token with an alternative sampled from the current student. This creates a comparison between two token choices in the same context. We then select one of these choices using the reference policy, complete the selected branch with that policy, and evaluate the resulting answer using the source reward oracle. A reward-dependent acceptance rule turns this observation into a statistical signal for estimating the calibration update. Repeating this construction at different token positions gathers information about the calibration coefficients across the generation process. The calibrated teacher subsequently provides supervision for student distillation on the target prompts. Under our model, the calibration learned from source comparisons also applies to target questions. The student update uses a batch of sampled answers and policy likelihoods to estimate its gradient and take one projected step. We compare this proposal with a randomly sampled student. After evaluating these two candidates, the selected student supplies token alternatives for the next calibration round. This coupled procedure transfers source reward information into target-side training through the calibrated teacher. The full pseudocode is given in Algorithm 1. CCL (Algorithm 1) maintains two trainable parameters: wt for the calibrated teacher πcal,t = πwt , and θt for the student πstu,θt . Line 1 initializes the calibration from the teacher. Each round first uses source rewards and a current-student alternative to update wt , then uses the updated calibrated teacher to train and select the next student. The frozen πtea supplies the source proposals, and πpre supplies the reference-weighted branch and completion. To obtain information for the calibration update, Lines 3-7 construct a token-level comparison. At a randomly selected position in a teacher answer, we retain the teacher prefix and compare its next token with a current-student alternative: st = (xit , at,1:ht −1 ), ct,1 = at,ht , ct,0 ∼ πstu,θt (· | st ). This comparison provides information about the calibration coefficients because the softmax parameterization gives zt := ϕ(st , ct,1 ) − ϕ(st , ct,0 ), log
πw (ct,1 | st ) = zt⊤ w. πw (ct,0 | st )
Thus, learning the optimal relative probabilities of these two tokens constrains the coefficients along zt . Sampling the branching position across all H steps collects the joint source information in (3.7). The next step uses source reward feedback to obtain a label for this comparison. Lines 8-13 select
9
Algorithm 1 Coupled Calibration and Learning (CCL) via Token-Level Branching Require: Source and target datasets {xi }ni=1 , {e x j }m j =1 ; source reward oracle R; teacher πtea = πwtea and πpre ; features ϕ, ϕstu and sets W , Θ; wtea ∈ W , θ0 ∈ Θ; λ, γ > 0, integer T ≥ 1; schedules (4.3). All draws use fresh randomness conditional on the preceding variables. 1: w0 ← wtea , πcal,0 ← πw0 . 2: for t = 0, . . . , T − 1 do Token-Level branching on source 3: Draw independently it ∼ Unif{1, . . . , n} and ht ∼ Unif{1, . . . , H}. 4: Draw at,1:H ∼ πtea (· | xit ) autoregressively. ▷ Teacher prefix and token. 5: st ← (xit , at,1:ht −1 ), ct,1 ← at,ht . 6: Draw ct,0 ∼ πstu,θt (· | st ). ▷ Student alternative. 7: zt ← ϕ(st , ct,1 ) − ϕ(st , ct,0 ). Teacher calibration. πpre (ct,1 | st ) . 8: Yt ∼ Bernoulli πpre (ct,1 | st ) + πpre (ct,0 | st ) pre pre ▷ Construct the selected branch. 9: at,1:ht −1 ← at,1:ht −1 , at,ht ← ct,Yt . pre pre 10: Draw at,ht +1:H ∼ πpre (· | xit , at,1:ht ) autoregressively. 11: Rt ← R(xit , apre ▷ Exactly one source reward query. t,1:H ). 12: Draw Ut ∼ Unif [0, 1]. 13: It ← 1{Ut ≤ exp((Rt − 1)/λ)}. 14: gt ← It zt [σ (zt⊤ wt ) − Yt ]. ▷ One trial; rejection gives gt = 0. 15: wt+1 ← ProjW (wt − ηt gt ). 16: πcal,t+1 ← πwt+1 . Student gradient update. π (a |x) 17: Define Zt+1 (θ, x, a1:H ) := λ log π stu,θ (a1:H |x) , Sstu,θ (x, a1:H ) := ∇θ log πstu,θ (a1:H | x), ∀θ ∈ Θ. cal,t+1 1:H 18: Draw independent prompt-answer pairs for ℓ = 1, . . . , bt+1 : jtg+1,ℓ ∼ Unif{1, . . . , m},
agt+1,ℓ,1:H | jtg+1,ℓ ∼ πstu,θt (· | x ej g
t+1,ℓ
).
bt + 1
1 X
, agt+1,ℓ,1:H )Zt+1 (θt ; x ej g , agt+1,ℓ,1:H ). t+1,ℓ ℓ=1 stu 20: ϑt+1,1 ← ProjΘ θt − αt+1 gbt+1 ; independently draw ϑt+1,2 ∼ Unif (Θ). Candidate student evaluation and selection. 21: for k = 1, 2 do 22: Draw fresh independent prompt-answer pairs for ℓ = 1, . . . , qt+1 : 19:
gbtstu +1 ←
bt + 1
Sstu,θt (x ej g
t+1,ℓ
jt+1,k,ℓ ∼ Unif{1, . . . , m}, 23:
bt+1,k ← C
1 qt+1
aval ejt+1,k,ℓ ). t+1,k,ℓ,1:H | jt+1,k,ℓ ∼ πstu,ϑt+1,k (· | x
qt + 1
X
Zt+1 (ϑt+1,k ; x ejt+1,k,ℓ , aval t+1,k,ℓ,1:H ).
ℓ=1
end for b bt+1,k , θt+1 ← ϑ kt+1 ← min argmink∈{1,2} C . t+1,b kt + 1 26: end for 27: return πstu,θT . 24: 25:
one reference-weighted branch, complete it with πpre , and query its reward Rt . The acceptance indicator satisfies Pr(It = 1 | Ft , st , ct,1 , ct,0 , Yt , Rt ) = exp((Rt − 1)/λ). Here Ft denotes the history before round t. The algorithm samples the branch and completion using the reference policy, and determines acceptance using the observed source reward. Under
10
Assumption 3.2, Lemma B.2 shows that the resulting accepted branch index satisfies πλ⋆ (ct,1 | st ) = σ (zt⊤ wλ⋆ ). πλ⋆ (ct,1 | st ) + πλ⋆ (ct,0 | st )
Pr(Yt = 1 | It = 1, Ft , st , ct,1 , ct,0 ) =
This identity characterizes the unknown parameter underlying the accepted labels. Algorithm 1 learns this parameter from the sampled labels by evaluating the logistic gradient at the current i h estimate gt = It zt σ (zt⊤ wt ) − Yt . Thus, source reward feedback provides logistic supervision for estimating wλ⋆ . This observation motivates the calibration update in Lines 14-15: gt = It zt [σ (zt⊤ wt ) − Yt ], wt+1 = ProjW (wt − ηt gt ), πcal,t+1 = πwt+1 . Here gt is the stochastic gradient of the acceptance-weighted logistic loss. The projection keeps the updated calibration parameter in W . The calibrated teacher then defines the student’s training objective on the target prompts: Ct + 1 ( θ ) : =
m λ X ej ) ∥ πwt+1 (· | x ej )) . KL(πstu,θ (· | x m j =1
(4.1)
Assumption 3.2 links source calibration to this target objective: the same wλ⋆ represents the regularized optimal policy on both datasets. The source-identification condition in (3.7) makes this parameter identifiable from source comparisons, and the known feature map evaluates the learned policy at target prefixes. Moreover, Lemma B.1 gives arg min Cwλ⋆ (θ ) = arg max Jλ,m (πstu,θ ). θ∈Θ
θ∈Θ
Thus, the calibration aligns the distillation objective with the regularized oracle-student objective. To update the student, Lines 18-19 estimate the gradient of Ct+1 using bt+1 independent currentstudent rollouts. The required quantities are the full-answer score and the sampled log-ratio cost: H X
Sstu,θ (x, a1:H ) := ∇θ log πstu,θ (a1:H | x) =
ϕstu (sh , ah ) −
h=1
Zt+1 (ϑ; x, a1:H ) := λ log
X
πstu,θ (b | sh )ϕstu (sh , b) ,
b∈B (sh )
H X πstu,ϑ (a1:H | x) πstu,ϑ (ah | x, a1:h−1 ) =λ log , πwt+1 (a1:H | x) πwt+1 (ah | x, a1:h−1 ) h=1
where sh = (x, a1:h−1 ) and B (sh ) is the legal-token set defined in the model setup. The score measures how the answer’s log probability changes with the student parameters. The cost Zt+1 measures its log-likelihood discrepancy from the calibrated teacher. For each ℓ ∈ {1, . . . , bt+1 }, we draw a fresh uniform target index jtg+1,ℓ and then sample an answer agt+1,ℓ,1:H from the current student at that question. The bt+1 prompt-answer pairs are independent conditional on the history and the calibration update. Averaging their score-cost products gives gbtstu +1 =
bt + 1 1 X
bt+1 ℓ=1
ej g Sstu,θt (x
t+1,ℓ
ej g , agt+1,ℓ,1:H )Zt+1 (θt ; x
t+1,ℓ
11
, agt+1,ℓ,1:H ).
Lemma B.6 shows that
h
i
E gbtstu +1 wt+1 , θt = ∇θ Ct+1 (θ )|θ =θt . The calibrated teacher is held fixed throughout this batch. Averaging preserves unbiasedness and reduces the conditional gradient variance by a factor of bt+1 . Line 20 uses this estimate to form two student candidates:
ϑt+1,1 = ProjΘ θt − αt+1 gbtstu +1 , ϑt+1,2 ∼ Unif ( Θ ). The projected-gradient candidate takes one step using the averaged gradient, with a constant step size justified by the smoothness bound in Lemma B.7. The random candidate explores the full parameter ball uniformly. In Lemma B.9, exploration provides progress outside a fixed neighborhood of the oracle student, while the gradient candidate provides quantitative progress near the oracle student. Each round thus constructs one projected-gradient proposal and the batch improves the accuracy of this update. To compare these candidates, the validation steps ending at Line 23 average Zt+1 over qt+1 fresh rollouts from each candidate, k ∈ {1, 2}, by Monte Carlo simulation: Cbt+1,k =
qt+1 1 X
qt + 1 ℓ = 1
ejt+1,k,ℓ , aval Zt+1 ϑt+1,k ; x t+1,k,ℓ,1:H .
Since the target indices are uniform and each answer is sampled from the candidate being evaluated, m i 1 X b Ea1:H ∼πstu,ϑ E Ct+1,k wt+1 , ϑt+1,k =
h
m j =1
t+1,k
ej , a1:H )] = Ct+1 (ϑt+1,k ). (·|e xj ) [Zt+1 (ϑt+1,k ; x
The sampled token log ratios therefore estimate the KL in the required student-to-calibrated-teacher direction. The algorithm selects the smallest estimated cost: kbt+1 = min argmin Cbt+1,k , θt+1 = ϑt+1,bk
t+1
k∈{1,2}
.
The same quantity Zt+1 consequently serves as both the weight in the student gradient estimate and the validation cost. Both quantities are computed from policy likelihoods and target rollouts, so target reward feedback is not required. The selected student supplies the alternative-token distribution in the next source comparison, completing the coupling between teacher calibration and student distillation. For the calibration step size and its analysis, let σ (u) = (1 + e−u )−1 and define γ := e−1/λ e−2B σ ′ (2B )µjoint > 0.
(4.2)
This constant quantifies the source-calibration information used in the convergence analysis. The algorithm uses γ to set its calibration step size. Define Lst := λH (1 + 8BH ), αstu :=
1 . 2Lst
Lemma B.7 proves that Lst is a uniform smoothness bound for the student objectives. For t =
12
0, . . . , T − 1, we use ηt =
1 , αt+1 = αstu , bt+1 = t + 2, qt+1 = (t + 2)2 , γ (t + 2)
(4.3)
where γ = e−1/λ e−2B σ ′ (2B )µjoint > 0 is defined in (4.2). These schedules specify the calibration step size, student step size, gradient batch size, and validation budget per candidate, respectively. The increasing gradient batch makes its sampling error vanish while the student step size remains fixed. All draws use fresh randomness conditional on their stated sampling laws, and the two validation batches are independent conditional on the history and candidate parameters. Thus target update t + 1 uses bt+1 + 2qt+1 full-answer rollouts: bt+1 for its gradient estimate and qt+1 for each candidate evaluation. Over T rounds, the algorithm makes exactly T source reward queries and uses O (T 3 ) target rollouts, of which O (T 2 ) estimate student gradients. Each round forms one projected student-gradient candidate.
5
Theoretical guarantee
In this section, we establish a finite-iteration convergence guarantee for CCL (Algorithm 1). We bound the average KL divergence between the returned student and the oracle student, accounting for teacher calibration, stochastic student gradients, and finite-rollout evaluation of the two student candidates. Throughout this section, the source and target datasets are fixed, and expectations include the complete adaptive randomness of the algorithm. Theorem 5.1. Under the model and assumptions in Section 3, there exist fixed-model constants Cλ,m > 0 and pλ,m ≥ 1, independent of T , such that the output of Algorithm 1 satisfies, for every integer T ≥ 1,
E
m 1 X
m j =1
ej ) πstu,θ† (· | x ej ) ≤ Cλ,m (T + 1)−1/(4pλ,m ) . KL πstu,θT (· | x
(5.1)
λ,m
Proof sketch. We first identify the objective that the student should minimize. The true regularized return equals a constant minus Cwλ⋆ (θ ). Hence, its minimizer over Θ is exactly the oracle † student. We study the excess true cost ∆λ,m (θ ) = Cwλ⋆ (θ ) − Cwλ⋆ (θλ,m ).
We next control the source calibration update. The accepted branch index is a logistic observation with parameter wλ⋆ . Source identification then gives h
i
E ∥wt − wλ⋆ ∥22 ≤
4 , t ≥ 1. γ 2 (t + 2)
This bound holds for the adaptive teacher sequence generated by CCL. For the student update, we first analyze a population projected gradient step for Cwλ⋆ . The softmax model gives a uniform smoothness bound. Analyticity and the uniqueness of the oracle yield a local Lojasiewicz inequality, which implies that this gradient step decreases the excess cost by at least a positive constant times its square in a fixed near-optimal region. Outside that region, the uniform candidate has a fixed positive probability of proposing a student with lower true cost. The population gradient step never increases the true cost, so selecting the better of the gradient
13
and uniform proposals combines these two sources of progress. Thus gradient descent supplies the improvement near the oracle, while exploration supplies progress outside that region. The algorithm uses a sampled gradient of Cwt instead of the population gradient of Cwλ⋆ . We control this difference using the minibatch variance and the calibration bound above. Fresh validation controls the error when comparing candidates. With et := E[∆λ,m (θt )], the resulting recursion is et ≤ et−1 − κopt e2t−1 +
16αstu λ2 B 2 H 3 16αstu λ2 BH 3 + 8λH 16λBH √ √ + + , t+1 t+1 γ t+2 | {z }
|
{z
}
gradient estimation
|
{z
calibration
}
validation
where κopt > 0 is independent of t. An induction then gives eT ≤ Kopt (T + 1)−1/4 for a fixed constant Kopt > 0. Finally, recall the average oracle-policy KL Kλ,m from (3.6). Its zero set contains the zero set of ∆λ,m , and both functions are analytic on the compact parameter ball. A second Lojasiewicz inequality therefore gives ∆λ,m (θ ) ≥ aλ,m [Kλ,m (θ )]pλ,m , aλ,m > 0, pλ,m ≥ 1. Applying Jensen’s inequality, we obtain E[Kλ,m (θT )] ≤
E[∆λ,m (θT )] aλ,m
!1/pλ,m
≤
Kopt aλ,m
!1/pλ,m
(T + 1)−1/(4pλ,m ) .
This proves the stated rate with Cλ,m = (Kopt /aλ,m )1/pλ,m . Appendix B gives the constants and the complete proof. Theorem 5.1 establishes that LLM distillation can recover the optimal student from a biased teacher, even when true reward feedback is entirely unavailable on the target dataset. Under some regularity assumptions, the expected average KL distance to the oracle student vanishes at a polynomial rate. Crucially, this oracle is defined by the true reference-regularized target objective, so the guarantee concerns the student’s actual target performance rather than its agreement with the teacher. The benchmark also respects the student’s limited capacity: the student class need not represent the unrestricted optimal policy. The result, therefore, connects source reward feedback to optimal target-side learning within a fixed student class. Through coupled calibration and distillation, source supervision corrects the policy that guides target training, allowing the student to overcome errors in the original teacher. Consequently, teacher bias need not impose a persistent loss relative to the best achievable student, and recovering this benchmark does not require collecting reward feedback on the target questions.
6
Separation from Direct Teacher Matching
In this section, to illustrate the power of our algorithm, we show that regularized direct matching can retain a positive error relative to the oracle student. Specifically, we construct a target instance in which the teacher is better than any model in the student model class; yet, direct matching learns a student policy that is strictly separated from the oracle student. The distillation problem that we consider is a LLM-as-judge setting. Fix H = 2, 1 ≤ d < D, λ > 0,
14
and α ∈ [1/2, 1). Each prompt contains a question and a candidate answer to be assessed. At h = 1, the model outputs a verdict: 1 declares the candidate correct, 0 declares it incorrect, and null expresses abstention due to uncertainty. Then when h = 2, it outputs EOS to terminate the answer. This setting models the practical task of distilling a compact LLM judge for automatic response evaluation. Training such judges on feedback from stronger models has been demonstrated using GPT-4-generated judgments and feedback (Zhu et al., 2025; Kim et al., 2024b). Our formulation captures the verdict-generation component of this task, including an explicit abstention option. ei,+ , x ei,− : i ∈ [d]}. For The fixed target dataset consists of m = 2d distinct prompts Dtar = {x e2i−1 := x ei,+ and x e2i := x ei,− . Let X = Dtar and an equivalent single-index enumeration, set x A = {0, 1, null, EOS}. The legal-token sets are
B ((x, ∅)) = {0, 1, null}, B ((x, a1 )) = {EOS}, a1 ∈ {0, 1, null}. The state transition appends the selected token, and the state (x, a1 , EOS) is terminal. Thus A(x) = {(0, EOS), (1, EOS), (null, EOS)}.
(6.1)
Here null is a first-step abstention verdict; the legal-token sets above specify this task’s output format. Every policy therefore satisfies π (EOS | x, a1 ) = 1, π ((a1 , EOS) | x) = π (a1 | x, ∅), a1 ∈ {0, 1, null}. ei,+ and Now, we model the ground truth verifier. We set that the candidate answer is correct at x ei,− . Hence, we define its ground-truth verdict and the exact evaluation reward by incorrect at x ei,+ ) = 1, y (x ei,− ) = 0, R(x, (a1 , EOS)) := 1{a1 = y (x)}. y (x
(6.2)
That is, a correct judgment receives reward 1. An incorrect judgment or abstention receives reward 0. In particular, the verdict 0 receives reward 1 when the candidate answer is incorrect. These target labels and rewards define the true benchmark; direct matching has access only to the target prompts and policy likelihoods. We assume that our pre-trained reference policy πpre is uniform over {0, 1, null}, πpre (a1 | x, ∅) =
1 , πpre (EOS | x, a1 ) = 1, a1 ∈ {0, 1, null}. 3
Let u1 , . . . , uD be an orthonormal basis of RD , and let e1 , . . . , ed be the standard basis of Rd . For the linear softmax teacher class, the feature vector is specified a =√1 a =√0 a = null ei,+ , ∅), a) u1 / 2 u2 / 2 ϕ((x 0 √ √ e ϕ((xi,− , ∅), a) u2 / 2 u1 / 2 0
(6.3)
For ϵ ∈ {+, −}, we define the student features by ei,ϵ , ∅), a) = ei 1{a = 1}, a ∈ {0, 1, null}. ϕstu ((x
(6.4)
For student and teacher classes, we set both feature maps to zero at every second-step state, and 15
set all unspecified feature values to zero. Every feature has norm at most one. Set √ d+2 B := , W := {w ∈ RD : ∥w∥2 ≤ B}, Θ := {θ ∈ Rd : ∥θ∥2 ≤ B}. λ We use the linear-softmax policies exp(θ⊤ ϕstu (s, a)) exp(w⊤ ϕ(s, a)) P , π ( a | s ) = . stu,θ ⊤ ⊤ c∈B (s) exp(w ϕ(s, c)) c∈B (s) exp(θ ϕstu (s, c))
πw (a | s) = P
(6.5)
At the second step, the single legal token EOS has probability e0 /e0 = 1. We set that our frozen teacher has parameter √ α 2 (6.6) wtea := u1 ∈ W , πtea := πwtea . λ The objective function in this post-training process is
2 m X ej , a1:h−1 ) 1 X π (ah | x . ej , a1:2 ) − λ Jλ,m (π ) := log Ea1:2 ∼π (·|exj ) R(x ej , a1:h−1 ) m j =1 π ( a | x pre h h=1
(6.7)
† The oracle student policy has the parameter θλ,m ∈ argmaxθ∈Θ Jλ,m (πstu,θ ).
At each target prompt, we use πλ⋆ to denote the unrestricted maximizer of the corresponding reward-minus-reference-KL objective over all laws in ∆(A(x)). † In the next proposition, we compute the optimal parameter θλ,m for the oracle student policy explicitly. It also verifies that the teacher is stronger than this oracle student, although the teacher itself is biased. † 1 Proposition 6.1. For the setting above, the unique oracle parameter is θλ,m = 4λ 1d ∈ int(Θ).
For every i ∈ [d], ϵ ∈ {+, −}, and feasible answer, exp(1{a1 = 1}/(4λ)) . (6.8) λ,m e1/(4λ) + 2 √ The unrestricted optimal policy is πλ⋆ = πwλ⋆ with wλ⋆ = ( 2/λ)u1 ∈ W . The student class cannot represent it, and Jλ,m (πtea ) > Jλ,m (πstu,θ† ). (6.9) ei,ϵ ) = πstu,θ† ((a1 , EOS) | x
λ,m
However, in direct matching distillation, we do not have the golden answers to the prompts, and thus we do not have access to the reward verifier function R. Therefore, we train our student based on the synthetic outputs from the fixed teacher model πtea . The student maximizes the teacher-student log-likelihood ratio minus the reference penalty: m ej ) ej ) πstu,θ (a1:2 | x 1 X πtea (a1:2 | x JSM (θ ) := E log − λ log . ej ) ej ) m j =1 a1:2 ∼πstu,θ (·|exj ) πstu,θ (a1:2 | x πpre (a1:2 | x
"
#
16
(6.10)
Equivalently, it minimizes the cost m 1 X ej ) ∥ πtea (· | x ej )) + λKL(πstu,θ (· | x ej ) ∥ πpre (· | x ej )) . CSM (θ ) := −JSM (θ ) = KL(πstu,θ (· | x m j =1
(6.11) For a feasible answer a1:2 , we define the evaluable cost and student score: πstu,θ (a1:2 | x) π (a1:2 | x) + λ log stu,θ , Sstu,θ (x, a1:2 ) := ∇θ log πstu,θ (a1:2 | x). πtea (a1:2 | x) πpre (a1:2 | x) (6.12) For the step size in the direct matching distillation algorithm, we set κSM = σ ′ (B + log 2) > 0, GSM = (2 + λ)B, µSM = (1+λd)κSM , ηtSM = µ (1t+2) . The pseudocode for the regularized SM direct-matching baseline is given in Algorithm 2.
ZSM (θ; x, a1:2 ) := log
Algorithm 2 Direct Teacher Matching SM Require: Fixed target prompts {e x j }m j =1 ; frozen πtea , πpre ; student class Θ; λ > 0; stepsizes ηt ; T ≥ 1. SM 1: θ0 ← 0. 2: for t = 0, . . . , T − 1 do 3: Draw jtSM ∼ Unif{1, . . . , m}. 4: Draw aSM ej SM ) autoregressively. t,1:2 ∼ πstu,θSM (· | x t
t
SM gbtSM ← Sstu,θSM (x ej SM , aSM ej SM , aSM t,1:2 )ZSM (θt ; x t,1:2 ).
5:
t
t
SM SM SM θtSM bt ). +1 ← ProjΘ (θt − ηt g
▷ Estimate the cost gradient.
t
6: 7: end for 8: return πstu,θSM . T
Now, we compare the output of direct matching with the explicit oracle student in Proposition 6.1. We will show that there is a separation between the trained student and the oracle student policy. Specifically, our next theorem shows that their KL divergence is lower bounded by a constant. Theorem 6.2. For the target setting in Section 6, for every T ≥ 1, Algorithm 2 satisfies
m κSM 1 X ej ) πstu,θ† (· | x ej ) ≥ E KL πstu,θSM (· | x T λ,m m j =1 2d
"√
#2
d(1 − α ) GSM √ − . 4λ µSM T + 1 +
Here [u]+ := max{u, 0}. In particular, for every integer T ≥
64λ2 dG2SM (1+λ)2 κ2SM (1−α)2
(6.13)
, the expected
average KL divergence is at least κSM (1 − α)2 /(128λ2 ). Therefore, we conclude that direct matching retains a nonvanishing distillation error relative to the oracle student πstu,θ† . λ,m
This theorem highlights the importance of teacher calibration in CCL. Regularized direct matching can retain teacher bias in the learned student, whereas CCL iteratively calibrates the policy that defines the student’s training objective. It complements Theorem 5.1: CCL converges under the source-calibration assumptions, whereas teacher quality alone does not guarantee that direct imitation recovers the oracle student.
17
7
Discussion
We developed Coupled Calibration and Learning (CCL) and a statistical framework for mitigating teacher bias in LLM distillation when reward feedback is available only for source questions. CCL couples teacher calibration with student updates through token-level branching, using source feedback to guide distillation on target questions. Each student update selects between a minibatch projected-gradient proposal and a uniformly sampled candidate using fresh target rollouts. We established its polynomial convergence in expected average KL divergence to the oracle student for the reference-regularized target objective. The proof quantifies progress from the student gradient update near the oracle and uses exploration to obtain progress outside a fixed near-optimal region. It also controls the errors from teacher calibration and finite-rollout estimation. This benchmark accounts for the limitations of the student class and does not require it to represent the unrestricted optimal policy. We also established a separation from regularized direct matching, which can retain a nonvanishing error even when the teacher outperforms every student policy. Together, these results identify coupled calibration and learning as a mechanism for recovering the optimal student without target-domain reward feedback. Several directions remain for future investigation. First, computational experiments with pretrained LLMs on coding and mathematical reasoning tasks would help assess the method under realistic rollout budgets, reward-evaluation costs, and source-target shifts. Such experiments could also examine the contributions of calibration, student updates, and candidate selection to practical performance. Second, our algorithm updates the calibration parameter and the student in every iteration. It would be useful to study less frequent calibration, with multiple student updates between successive calibration steps. The central question is whether such schedules preserve convergence to the oracle student while reducing calibration costs, and how their relative update frequencies affect the rate. Finally, our analysis gives an explicit one-quarter exponent for the excess true cost, while conversion to oracle-policy KL still uses a model-dependent Lojasiewicz exponent. The constants also depend on the local geometry and the probability of exploring a fixed near-optimal region. Deriving explicit bounds on these quantities is an important theoretical direction. A sharper analysis may exploit the stronger local gradient inequality before its quadratic relaxation, reduce the gradient and validation budgets, and clarify the dependence on the horizon and model dimensions.
References Alekh Agarwal, Sham M. Kakade, Jason D. Lee, and Gaurav Mahajan. On the theory of policy gradient methods: Optimality, approximation, and distribution shift. Journal of Machine Learning Research, 22(98):1–76, 2021. Anastasios N. Angelopoulos, Stephen Bates, Clara Fannjiang, Michael I. Jordan, and Tijana Zrnic. Prediction-powered inference. Science, 382(6671):669–674, 2023. Yu Bai, Fan Chen, Huan Wang, Caiming Xiong, and Song Mei. Transformers as statisticians: Provable in-context learning with in-context algorithm selection. In Advances in Neural Information Processing Systems, volume 36, pages 57125–57211, 2023. Peter L Bartlett, Philip M Long, Gábor Lugosi, and Alexander Tsigler. Benign overfitting in linear regression. Proceedings of the National Academy of Sciences, 117(48):30063–30070, 2020. 18
Edward Bierstone and Pierre D. Milman. Mathématiques de l’IHÉS, 67:5–42, 1988.
Semianalytic and subanalytic sets.
Publications
Jérôme Bolte, Aris Daniilidis, and Adrian Lewis. The lojasiewicz inequality for nonsmooth subanalytic functions with applications to subgradient dynamical systems. SIAM Journal on Optimization, 17(4):1205–1223, 2007. Fan Chen, Zeyu Jia, Alexander Rakhlin, and Tengyang Xie. Outcome-based online reinforcement learning: Algorithms and fundamental limits. In Advances in Neural Information Processing Systems, volume 38, 2025. Fan Chen, Audrey Huang, Noah Golowich, Sadhika Malladi, Adam Block, Jordan T. Ash, Akshay Krishnamurthy, and Dylan J. Foster. The coverage principle: How pre-training enables posttraining. In The Fourteenth International Conference on Learning Representations, 2026. Mark Chen, Jerry Tworek, Heewoo Jun, Qiming Yuan, Henrique Ponde de Oliveira Pinto, et al. Evaluating large language models trained on code, 2021. arXiv preprint arXiv:2107.03374. Wojciech M. Czarnecki, Razvan Pascanu, Simon Osindero, Siddhant Jayakumar, Grzegorz Swirszcz, and Max Jaderberg. Distilling policy distillation. In Proceedings of the Twenty-Second International Conference on Artificial Intelligence and Statistics, volume 89 of Proceedings of Machine Learning Research, pages 1331–1340, 2019. Tri Dao, Govinda M. Kamath, Vasilis Syrgkanis, and Lester Mackey. Knowledge distillation as semiparametric inference. In International Conference on Learning Representations, 2021. DeepSeek-AI. DeepSeek-R1: Incentivizing reasoning capability in LLMs via reinforcement learning, 2025. arXiv preprint arXiv:2501.12948. Dylan J. Foster, Adam Block, and Dipendra Misra. Is behavior cloning all you need? understanding horizon in imitation learning. In Advances in Neural Information Processing Systems, volume 37, pages 120602–120666, 2024. Dylan J. Foster, Zakaria Mhammedi, and Dhruv Rohatgi. Is a good foundation necessary for efficient reinforcement learning? the computational role of the base model in exploration. In Proceedings of Thirty Eighth Conference on Learning Theory, volume 291 of Proceedings of Machine Learning Research, pages 2026–2142, 2025. Jiawei Ge, Shange Tang, Jianqing Fan, Cong Ma, and Chi Jin. Maximum likelihood estimation is all you need for well-specified covariate shift. In The Twelfth International Conference on Learning Representations, 2024. Gemini Team. Gemini: A family of highly capable multimodal models. arXiv:2312.11805, 2023.
arXiv preprint
Aaron Grattafiori, Abhimanyu Dubey, Abhinav Jauhri, et al. The Llama 3 herd of models. arXiv preprint arXiv:2407.21783, 2024. Yuxian Gu, Li Dong, Furu Wei, and Minlie Huang. MiniLLM: Knowledge distillation of large language models. In The Twelfth International Conference on Learning Representations, 2024.
19
Srishti Gureja, Lester James V. Miranda, Shayekh Bin Islam, Rishabh Maheshwary, Drishti Sharma, Gusti Winata, Nathan Lambert, Sebastian Ruder, Sara Hooker, and Marzieh Fadaee. MRewardBench: Evaluating reward models in multilingual settings. In Proceedings of the 63rd Annual Meeting of the Association for Computational Linguistics (Volume 1: Long Papers), pages 43–58, 2025. Geoffrey Hinton, Oriol Vinyals, and Jeff Dean. Distilling the knowledge in a neural network, 2015. arXiv preprint arXiv:1503.02531. Haichen Hu and David Simchi-Levi. Perturbing the derivative: Doubly wild refitting for model-free evaluation of opaque machine learning predictors, 2025a. Haichen Hu and David Simchi-Levi. Perturbing the derivative: Wild refitting for model-free evaluation of machine learning models under bregman losses, 2025b. Haichen Hu and David Simchi-Levi. Interleaved resampling and refitting: Data and computeefficient evaluation of black-box predictors. arXiv preprint arXiv:2603.14218, 2026. Haichen Hu, Jian Qian, and David Simchi-Levi. Model-based reinforcement learning with double oracle efficiency in policy optimization and offline estimation, 2026. Audrey Huang, Adam Block, Dylan J. Foster, Dhruv Rohatgi, Cyril Zhang, Max Simchowitz, Jordan T. Ash, and Akshay Krishnamurthy. Self-improvement in language models: The sharpening mechanism. In The Thirteenth International Conference on Learning Representations, 2025. Muhammed Emrullah Ildiz, Halil Alperen Gozeten, Ege Onur Taga, Marco Mondelli, and Samet Oymak. High-dimensional analysis of knowledge distillation: Weak-to-strong generalization and scaling laws. In The Thirteenth International Conference on Learning Representations, 2025. Fotis Iliopoulos, Vasilis Kontonis, Cenk Baykal, Gaurav Menghani, Khoa Trinh, and Erik Vee. Weighted distillation with unlabeled examples. In Advances in Neural Information Processing Systems, volume 35, 2022. Zeyu Jia, Alexander Rakhlin, and Tengyang Xie. Do we need to verify step by step? rethinking process supervision from a theoretical perspective. In Proceedings of the 42nd International Conference on Machine Learning, volume 267 of Proceedings of Machine Learning Research, pages 27373–27398, 2025. Nan Jiang, Akshay Krishnamurthy, Alekh Agarwal, John Langford, and Robert E. Schapire. Contextual decision processes with low Bellman rank are PAC-learnable. In Proceedings of the 34th International Conference on Machine Learning, volume 70 of Proceedings of Machine Learning Research, pages 1704–1713, 2017. Chi Jin, Zhuoran Yang, Zhaoran Wang, and Michael I Jordan. Provably efficient reinforcement learning with linear function approximation. In Conference on learning theory, pages 2137–2143. PMLR, 2020. Juno Kim, Tai Nakamaki, and Taiji Suzuki. Transformers are minimax optimal nonparametric in-context learners. In Advances in Neural Information Processing Systems, volume 37, pages 106667–106713, 2024a.
20
Juno Kim, Jihun Yun, Jason D. Lee, and Kwang-Sung Jun. Coverage improvement and fast convergence of on-policy preference learning. In Proceedings of the 43rd International Conference on Machine Learning, 2026. Seungone Kim, Jamin Shin, Yejin Cho, Joel Jang, Shayne Longpre, Hwaran Lee, Sangdoo Yun, Seongjin Shin, Sungdong Kim, James Thorne, and Minjoon Seo. Prometheus: Inducing finegrained evaluation capability in language models. In International Conference on Learning Representations, 2024b. Qi Lei, Wei Hu, and Jason D. Lee. Near-optimal linear regression under distribution shift. In Proceedings of the 38th International Conference on Machine Learning, volume 139 of Proceedings of Machine Learning Research, pages 6164–6174, 2021. Hunter Lightman, Vineet Kosaraju, Yuri Burda, Harrison Edwards, Bowen Baker, Teddy Lee, Jan Leike, John Schulman, Ilya Sutskever, and Karl Cobbe. Let’s verify step by step. In The Twelfth International Conference on Learning Representations, 2024. Michal Lukasik, Srinadh Bhojanapalli, Aditya Krishna Menon, and Sanjiv Kumar. Teacher’s pet: understanding and mitigating biases in distillation. Transactions on Machine Learning Research, 2022. Cong Ma, Reese Pathak, and Martin J. Wainwright. Optimally tackling covariate shift in RKHSbased nonparametric regression. The Annals of Statistics, 51(2):738–761, 2023. Aditya Krishna Menon, Ankit Singh Rawat, Sashank J. Reddi, Seungyeon Kim, and Sanjiv Kumar. A statistical perspective on distillation. In Proceedings of the 38th International Conference on Machine Learning, volume 139 of Proceedings of Machine Learning Research, pages 7632–7642, 2021. OpenAI. GPT-4 technical report, 2023. arXiv preprint arXiv:2303.08774. Jian Qian, Haichen Hu, and David Simchi-Levi. Offline oracle-efficient learning for contextual mdps via layerwise exploration-exploitation tradeoff. arXiv preprint arXiv:2405.17796, 2024. Qwen Team. Qwen2.5 technical report. arXiv preprint arXiv:2412.15115, 2024. Stephane Ross, Geoffrey Gordon, and Drew Bagnell. A reduction of imitation learning and structured prediction to no-regret online learning. In Proceedings of the Fourteenth International Conference on Artificial Intelligence and Statistics, volume 15 of Proceedings of Machine Learning Research, pages 627–635, 2011. Dmitri Roussinov, Serge Sharoff, and Nadezhda Puchnina. Controlling out-of-domain gaps in LLMs for genre classification and generated text detection. In Proceedings of the 31st International Conference on Computational Linguistics, pages 3329–3344, 2025. Nikunj Saunshi, Sadhika Malladi, and Sanjeev Arora. A mathematical exploration of why language models help solve downstream tasks. In International Conference on Learning Representations, 2021. Ved Sriraman and Adam Block. Revisiting the (sub)optimality of best-of-N for inference-time alignment. In Proceedings of Thirty Ninth Conference on Learning Theory, volume 336 of Proceedings of Machine Learning Research, pages 5980–6028, 2026. 21
Ved Sriraman, Peihan Liu, Daniel Hsu, and Adam Block. Behavior cloning is not all you need: The optimality of on-policy distillation for noisy expert feedback, 2026. Martin J Wainwright. Wild refitting for black box prediction. arXiv preprint arXiv:2506.21460, 2025. Kaizheng Wang. Pseudo-labeling for kernel ridge regression under covariate shift. The Annals of Statistics, 54(1):252–276, 2026. Nathan Weill and Kaizheng Wang. Pseudo-labeling for unsupervised domain adaptation with kernel GLMs, 2026. Eric Xia and Jason M. Klusowski. Classification imbalance as transfer learning, 2026. Eric Xia and Martin J. Wainwright. Prediction aided by surrogate training, 2024. Audrey Xie, Ludwig Schmidt, and John Duchi. Two mathematical models of knowledge distillation. In Proceedings of The 29th International Conference on Artificial Intelligence and Statistics, volume 300 of Proceedings of Machine Learning Research, pages 4726–4734, 2026. Tengyang Xie, Ching-An Cheng, Nan Jiang, Paul Mineiro, and Alekh Agarwal. Bellman-consistent pessimism for offline reinforcement learning. In Advances in Neural Information Processing Systems, volume 34, pages 6683–6694, 2021. Tengyang Xie, Dylan J. Foster, Yu Bai, Nan Jiang, and Sham M. Kakade. The role of coverage in online reinforcement learning. In The Eleventh International Conference on Learning Representations, 2023. Tengyang Xie, Dylan J. Foster, Akshay Krishnamurthy, Corby Rosset, Ahmed H. Awadallah, and Alexander Rakhlin. Exploratory preference optimization: Harnessing implicit Q*-approximation for sample-efficient RLHF. In The Thirteenth International Conference on Learning Representations, 2025. Wei Xiong, Hanze Dong, Chenlu Ye, Ziqi Wang, Han Zhong, Heng Ji, Nan Jiang, and Tong Zhang. Iterative preference learning from human feedback: Bridging theory and practice for RLHF under KL-constraint. In Proceedings of the 41st International Conference on Machine Learning, volume 235 of Proceedings of Machine Learning Research, pages 54715–54754, 2024. Kakei Yamamoto and Martin J. Wainwright. Residual-as-teacher: Mitigating bias propagation in student–teacher estimation, 2026. Chenlu Ye, Wei Xiong, Yuheng Zhang, Hanze Dong, Nan Jiang, and Tong Zhang. Online iterative reinforcement learning from human feedback with general preference model. In Advances in Neural Information Processing Systems, volume 37, 2024. Zhuohao Yu, Zhiwei Steven Wu, and Adam Block. From curiosity to caution: Mitigating reward hacking for best-of-N with pessimism. In The Fourteenth International Conference on Learning Representations, 2026. Yurun Yuan, Fan Chen, Zeyu Jia, Alexander Rakhlin, and Tengyang Xie. Trajectory Bellman residual minimization: A simple value-based method for LLM reasoning. In Advances in Neural Information Processing Systems, volume 38, 2025.
22
Huaqing Zhang, Jingchu Gai, Juno Kim, Bingbin Liu, and Andrej Risteski. When does online imitation learning help in LLM post-training? the role of (non-)realizability beyond horizon, 2026. Yuheng Zhang, Yu Bai, and Nan Jiang. Offline learning in Markov games with general function approximation. In Proceedings of the 40th International Conference on Machine Learning, volume 202 of Proceedings of Machine Learning Research, pages 40804–40829, 2023. Heyang Zhao, Chenlu Ye, Quanquan Gu, and Tong Zhang. Sharp analysis for KL-regularized contextual bandits and RLHF. In Advances in Neural Information Processing Systems, volume 38, 2025. Banghua Zhu, Michael I. Jordan, and Jiantao Jiao. Principled reinforcement learning with human feedback from pairwise or K-wise comparisons. In Proceedings of the 40th International Conference on Machine Learning, volume 202 of Proceedings of Machine Learning Research, pages 43037–43067, 2023. Lianghui Zhu, Xinggang Wang, and Xinlong Wang. JudgeLM: Fine-tuned large language models are scalable judges. In International Conference on Learning Representations, 2025.
A
Mathematical Tools
Lemma A.1 (Lojasiewicz inequality Bierstone and Milman (1988)). Let K ⊆ Rq be nonempty, and let f , g : K → R. Suppose that their graphs are compact subanalytic subsets of Rq +1 and that {x ∈ K : f (x) = 0} ⊆ {x ∈ K : g (x) = 0}. Then there are constants c > 0 and ρ > 0 for which |f (x)| ≥ c|g (x)|ρ (x ∈ K ). In particular, if f , g ≥ 0 on K, we choose any M ≥ max{1, supx∈K g (x)}. Then, the same constants c, ρ defined above yield f (x) ≥ a[g (x)]p , p := max{1, ρ} ≥ 1, a := cM ρ−p > 0. To quantify the progress of a projected student step near the oracle, we use the following subgradient form of the Lojasiewicz inequality. The domain in this result can include the boundary of the student parameter ball. Lemma A.2 (Bolte et al., 2007, Theorem 3.1). Let f : Rd → R ∪ {+∞} have a subanalytic graph and closed domain, and suppose that f is continuous on its domain. If θ̄ is a critical point, meaning 0 ∈ ∂f (θ̄ ), then there are a neighborhood U of θ̄, c > 0, and ρ ∈ [0, 1) such that dist(0, ∂f (θ )) ≥ c |f (θ ) − f (θ̄ )|ρ for θ ∈ U ∩ dom f with f (θ ) ̸= f (θ̄ ). Here ∂f is the limiting subdifferential and dist(0, ∂f (θ )) := inf v∈∂f (θ ) ∥v∥2 .
23
B
Proofs in Section 5
In this appendix, we provide the proof of Theorem 5.1. Specifically, we will first provide a sequence of lemmas that are useful and finally, we will show how to combine them together to prove the main convergence rate theorem. We begin by identifying the correct distillation objective: the following lemma shows that maximizing the regularized student return is equivalent to minimizing the KL cost against πwλ⋆ . Lemma B.1. For each source or target prompt x, define πpre (b1:H | x)eR(x,b1:H )/λ .
X
Zλ ( x ) : =
b1:H ∈A(x)
The unique unrestricted optimal answer policy is πpre (a1:H | x)eR(x,a1:H )/λ . Zλ ( x )
(B.1)
m λ X e j ) − Cw ⋆ ( θ ) . log Zλ (x Jλ,m (πstu,θ ) = λ m j =1
(B.2)
πλ⋆ (a1:H | x) = Moreover, for every student parameter,
The Gibbs representation in Lemma B.1 lets us identify the distribution of accepted branch indices. The following lemma shows that these indices form logistic observations with parameter wλ⋆ , providing the statistical basis for the source update. We now define a filtration {Ft }Tt=1 . We treat the source and target datasets, the reward oracle, the policy features, the frozen policies, and the algorithmic schedules as fixed. Let Ft denote the information available immediately before the source sampling in round t. Formally, set F0 = σ (w0 , θ0 ). For every t > 0, we define Ft recursively
g g Ft+1 = Ft ∨ σ it , ht , at,1:H , ct,0 , Yt , apre t,1:H , Ut , jt+1,ℓ , at+1,ℓ,1:H
bt+1 ℓ=1
, ϑt+1,2 , jt+1,k,ℓ , aval t+1,k,ℓ,1:H
k∈{1,2} ℓ=1,...,qt+1
The three lines collect the source samples, the student-gradient batch, and the exploration candidate and candidate-validation samples, respectively. All other quantities computed in round t, including st , ct,1 , zt , Rt , It , gt , the gradient candidate ϑt+1,1 , and the selected student, are measurable functions of Ft and these samples. In particular, wt and θt are Ft -measurable, whereas wt+1 and θt+1 are Ft+1 -measurable. For deterministic initialization, F0 is the trivial sigma-algebra. Throughout the following lemmas, conditioning on Ft , st , ct,1 , ct,0 means conditioning on Ft ∨ σ (st , ct,1 , ct,0 ). Then, we have the following lemma. Lemma B.2. Under Assumption 3.2, for every round t ≥ 0, Algorithm 1 satisfies, almost surely, Pr(It = 1 | Ft , st , ct,1 , ct,0 ) ≥ e−1/λ .
24
(B.3)
.
Moreover, conditional on acceptance, the branch index satisfies Pr(Yt = 1 | It = 1, Ft , st , ct,1 , ct,0 ) =
πλ⋆ (ct,1 | st ) = σ (zt⊤ wλ⋆ ), πλ⋆ (ct,1 | st ) + πλ⋆ (ct,0 | st )
(B.4)
where
1 . 1 + e−u These statements also hold when the two proposal tokens coincide, including prefixes after EOS. The random variable Yt denotes the branch index, not the token value. zt = ϕ(st , ct,1 ) − ϕ(st , ct,0 ), σ (u) =
Combining the accepted-label identity in Lemma B.2 with joint source identification, we obtain convergence of the calibration updates despite the adaptive student proposals. The bounds below will control both the calibration error and the changes in the target objective across rounds. Lemma B.3. With Gjoint and γ defined in (3.7)-(4.2), every round t ≥ 0 satisfies h
i
(B.5)
i
(B.6)
E zt zt⊤ | Ft ⪰ e−2B Gjoint , ∥zt ∥2 ≤ 2, h
E (wt − wλ⋆ )⊤ gt Ft ≥ γ∥wt − wλ⋆ ∥22 , ∥gt ∥2 ≤ 2. Consequently, for every integer t ≥ 1, the projected updates in Algorithm 1 satisfies 4 , γ 2 (t + 2) 2 ∥wt − wt−1 ∥2 ≤ . γ (t + 1)
h
i
E ∥wt − wλ⋆ ∥22 ≤
(B.7) (B.8)
Having controlled the source updates, we turn to the quantities estimated from target rollouts. We first bound the log-ratio cost of a single answer, which will be used to control the population cost and the error in its Monte Carlo evaluation. Lemma B.4. For t = 0, . . . , T , wt ∈ W , ϑ ∈ Θ, and a feasible answer a1:H ∈ A(x) at a target prompt x, define Zt (ϑ; x, a1:H ) := λ log
H X πstu,ϑ (a1:H | x) πstu,ϑ (ah | x, a1:h−1 ) =λ log . πwt (a1:H | x) πwt (ah | x, a1:h−1 ) h=1
Then, uniformly over these choices, |Zt (ϑ; x, a1:H )| ≤ 4λBH.
(B.9)
The preceding trajectory bound gives a uniform bound on its expected cost. We also quantify sensitivity to the student and calibration parameters, so that parameter changes can be translated into cost changes in the optimization analysis. λ Lemma B.5. For w ∈ W and θ ∈ Θ, define Cw (θ ) := m
25
ej ) ∥ πw (· | x ej )). j =1 KL(πstu,θ (· | x
ej , a1:H )], where Zt is j =1 Ea1:H ∼πstu,θ (·|e xj ) [Zt (θ; x
1 Pm
The cost used in round t is Ct (θ ) := Cwt (θ ) = m
Pm
defined in Lemma B.4. For every w, v ∈ W and θ, ϑ ∈ Θ, 0 ≤ Cw (θ ) ≤ 4λBH,
(B.10)
|Cw (θ ) − Cw (ϑ)| ≤ 4λBH 3/2 ∥θ − ϑ∥2 ,
(B.11)
sup |Cw (θ ) − Cv (θ )| ≤ 2λH∥w − v∥2 .
(B.12)
θ∈Θ
In particular, these bounds apply to Ct by setting w = wt . For the newly calibrated objective Ct+1 , we next verify that the algorithm’s fresh target batch provides an unbiased gradient estimate. We also bound its variance, which controls the error in the projected-gradient proposal. Lemma B.6. Let Ft be the history before the source draws in round t, so that θt is Ft -measurable, and set Gt+1 := Ft ∨ σ (wt+1 ). Define the score Sstu,θ (x, a1:H ) := ∇θ log πstu,θ (a1:H | x). Conditional on Gt+1 , independently for ℓ = 1, . . . , bt+1 , draw ej g jtg+1,ℓ ∼ Unif{1, . . . , m}, agt+1,ℓ,1:H ∼ πstu,θt (· | x
t+1,ℓ
),
and define ej g gbtstu +1,ℓ : = Sstu,θt (x
t+1,ℓ
ej g , agt+1,ℓ,1:H )Zt+1 (θt ; x
t+1,ℓ
, agt+1,ℓ,1:H ), gbtstu +1 : =
bt + 1 1 X
bt+1 ℓ=1
gbtstu +1,ℓ .
Here Zt+1 and Ct+1 = Cwt+1 are defined in Lemmas B.4 and B.5, respectively. For every t = 0, . . . , T − 1, almost surely, we have h
i
i
h
btstu E gbtstu +1 wt+1 , θt = ∇θ Ct+1 (θt ), +1 Gt+1 = E g
(B.13)
16λ2 B 2 H 3 . bt+1
(B.14)
E
gbtstu + 1 − ∇ θ Ct + 1 ( θ t )
2 2
Gt+1 ≤
The preceding bounds control the size of the student gradient. We next bound its variation, both as the student parameter changes and as the calibrated teacher changes. These bounds determine a fixed student step size and control the error caused by using the current calibration. Lemma B.7. Recall the cost Cw in (3.5) and the oracle-policy average KL divergence Kλ,m in P ej ), the function θ 7→ m−1 m (3.6). For any answer laws Qj that are positive on A(x j =1 KL(πstu,θ (· | d ej )∥Qj ) is real analytic on R . In particular, this holds for Cw and Kλ,m . Define x Lst := λH (1 + 8BH ), αstu :=
26
1 . 2Lst
For all w, v ∈ W and θ, ϑ ∈ Θ, we have ∥∇2θ Cw (θ )∥op ≤ Lst ,
(B.15)
∥∇θ Cw (θ ) − ∇θ Cw (ϑ)∥2 ≤ Lst ∥θ − ϑ∥2 , ∥∇θ Cw (θ ) − ∇θ Cv (θ )∥2 ≤ 2λH
3/2
∥w − v∥2 .
(B.16) (B.17)
After constructing the candidates, the algorithm chooses among them using sampled costs. The uniform log-ratio bound in Lemma B.4 allows the following lemma to control how far the selected candidate’s true cost can exceed the better of the two proposals. Lemma B.8. For each target update t ≥ 1, define Ht := Ft−1 ∨ σ (wt , ϑt,1 , ϑt,2 ). Define the nonnegative, proof-only comparison error ξt := 2 maxk∈{1,2} |Cbt,k − Ct (ϑt,k )|. Then, we have Ct (θt ) ≤ min Ct (ϑt,k ) + ξt , E[ξt | Ht ] ≤ k∈{1,2}
16λBH . t+1
(B.18)
The next lemma combines two sources of improvement for the fixed true cost. A projected gradient step provides progress near the oracle; uniform exploration provides progress when the current cost is bounded away from the minimum. The exact true gradient is used only to define the comparison step in this lemma. † Lemma B.9. Let F (θ ) := Cwλ⋆ (θ ) and ∆λ,m (θ ) := F (θ ) − F (θλ,m ). Define
y (θ ) := ProjΘ (θ − αstu ∇θ F (θ )) , ϑunif ∼ Unif (Θ), Mopt := 8λB 2 H 3/2 . Under Assumption 3.3, there is a fixed-model constant 0 < κopt ≤ 1/(4Mopt ) such that, for every θ ∈ Θ, h n oi Eϑunif F (θ ) − min F (y (θ )), F (ϑunif ) ≥ κopt [∆λ,m (θ )]2 . (B.19) The constant does not depend on the iteration number. We now compare the actual student update with the population update in Lemma B.9. The preceding gradient and selection bounds control the difference between these updates. This gives a recursion for the excess true cost, with separate contributions from gradient estimation, calibration, and validation. Lemma B.10. Recall ∆λ,m (θ ) and εt from (3.6). Let κopt > 0 be the constant in Lemma B.9, chosen so that κopt ≤ 1/(4Mopt ), where Mopt := 8λB 2 H 3/2 . Define Eopt := 16αstu λ2 B 2 H 3 + (
Kopt := max Mopt ,
s
16αstu λ2 BH 3 + 8λH + 16λBH, γ
2Eopt 1 , κopt 2κopt
)
.
Then Algorithm 1 satisfies, for every integer T ≥ 1, 8λH E[∆λ,m (θT )] ≤ Kopt (T + 1)−1/4 , E[εT ] ≤ Kopt (T + 1)−1/4 + √ . γ T +2
(B.20)
Lemma B.10 controls the expected excess true cost, whereas Theorem 5.1 concerns KL divergence 27
to the oracle student. Under uniqueness of the optimal student policy, the following lemma uses analyticity and compactness to connect these two errors. † Lemma B.11. Under the model setup and Assumption 3.3, fix θλ,m ∈ argminθ∈Θ Cwλ⋆ (θ ). Recall the excess true cost and the average oracle-policy KL from (3.6): † ∆λ,m (θ ) := Cwλ⋆ (θ ) − Cwλ⋆ (θλ,m ), m 1 X ej ) . ej ) πstu,θ† (· | x Kλ,m (θ ) := KL πstu,θ (· | x λ,m m j =1
There exist constants aλ,m > 0 and pλ,m ≥ 1 such that ∆λ,m (θ ) ≥ aλ,m [Kλ,m (θ )]pλ,m for every θ ∈ Θ.
(B.21)
These constants depend on the fixed model and target design, but not on the iteration number. Combining Lemmas B.10 and B.11 through Jensen’s inequality now yields the convergence rate in Theorem 5.1, as shown in the proof below. Proof of Theorem 5.1. Starting from the left-hand side of (5.1), the definition in (3.6) gives
m 1 X ej ) πstu,θ† (· | x ej ) = E[Kλ,m (θT )]. KL πstu,θT (· | x E λ,m m j =1
We first apply Lemma B.11 to compare this policy error with the excess true cost. Its inequality holds for every θ ∈ Θ, and hence on every sample path at θT : aλ,m [Kλ,m (θT )]pλ,m ≤ ∆λ,m (θT ). Dividing by aλ,m > 0 and taking the increasing power 1/pλ,m yields −1/p
Kλ,m (θT ) ≤ aλ,m λ,m [∆λ,m (θT )]1/pλ,m . The excess cost is bounded by Lemma B.5, so the expectations below are finite. Since pλ,m ≥ 1, the function u 7→ u1/pλ,m is concave: for u > 0, d2 1/pλ,m 1 u = 2 du pλ,m
1 pλ,m
!
− 1 u1/pλ,m −2 ≤ 0,
and continuity extends concavity to zero. Taking expectations and applying Jensen’s inequality therefore gives −1/p E[Kλ,m (θT )] ≤ aλ,m λ,m E [∆λ,m (θT )]1/pλ,m
h
i
−1/p ≤ aλ,m λ,m (E[∆λ,m (θT )])1/pλ,m =
E[∆λ,m (θT )] aλ,m
!1/pλ,m
Next, we use Lemma B.10 to bound the numerator by Kopt (T + 1)−1/4 . This bound already combines the student gradient progress, calibration error, and finite-rollout errors. Substituting it
28
.
into the increasing power gives E[Kλ,m (θT )] ≤
Kopt (T + 1)−1/4 aλ,m
!1/pλ,m
=
Kopt aλ,m
!1/pλ,m
(T + 1)−1/(4pλ,m ) .
Taking Cλ,m := (Kopt /aλ,m )1/pλ,m proves (5.1). All constants are independent of T , and pλ,m ≥ 1 is finite, so the bound converges to zero.
C
Proofs in Appendix B
In this appendix, we provide the proofs of the lemmas in Appendix B Proof of Lemma B.1. We fix x first and call the right-hand side of (B.1) as π̄λ (a1:H | x). Since R ∈ [0, 1] and the pre-trained policy πpre is a probability law, we have 1 ≤ eR(x,a1:H )/λ ≤ e1/λ =⇒ 1 ≤ Zλ (x) ≤ e1/λ . Thus, π̄λ is positive and sums to one. By the definition of an autoregressive policy, we have that π̄λ (ah | x, a1:h−1 ) =
π̄λ (a1:h | x) . π̄λ (a1:h−1 | x)
Thus, π̄λ is an admissible autoregressive policy. For any competing policy π, by the autoregressive product in (3.1) gives, we have that H X
log
h=1
Recall that π̄λ (a1:H |x) = we obtain
π (ah | x, a1:h−1 ) π (a1:H | x) = log . πpre (ah | x, a1:h−1 ) πpre (a1:H | x)
(C.1)
πpre (a1:H |x)eR(x,a1:H )/λ . Taking the logarithm on both sides and rearrange, Zλ ( x )
R(x, a1:H ) = λ log
π̄λ (a1:H | x) + λ log Zλ (x). πpre (a1:H | x)
(C.2)
Theerfore, we Subtract (C.1) from both sides of (C.2) and then take expectation to get "
Ea1:H ∼π (·|x)
H X
#
π (ah | x, a1:h−1 ) π (a1:H | x) R(x, a1:H ) − λ = λ log Zλ (x) − λEa1:H ∼π (·|x) log log π ( a | x, a ) π̄ pre h 1:h−1 λ (a1:H | x) h=1
= λ log Zλ (x) − λKL(π (· | x) ∥ π̄λ (· | x)) . The second is by the definition of KL. The KL divergence is nonnegative and KL(p∥p) = 0, Thus, the preceding objective is uniquely maximized by π = π̄λ , which proves (B.1). Average the identity over the target prompts and substitute πλ⋆ = πwλ⋆ from Assumption 3.2. This proves (B.2) and the equivalence of the original regularized benchmark and the minimum of Cwλ⋆ . The unknown target rewards appear in the analysis, not in an algorithmic query.
29
Let Ft be the history immediately before the source step of round t, containing all source and target random draws from rounds 0, . . . , t − 1 and the initial parameters. At t = 0, there are no preceding rounds. Thus wt and θt are Ft -measurable. The source draws in round t are fresh conditional on this history. In particular, the student may depend on every preceding calibration update, exploration draw, and target evaluation. Proof of Lemma B.2. We show that the acceptance rule reweights the two branch indices by exactly the factors appearing in the regularized optimal policy. Fix the history Ft , the selected prefix st = (xit , at,1:ht −1 ), and the two legal proposal tokens ct,1 and ct,0 . Here Ft is the history before round t. After additionally conditioning on Yt = k, the first ht tokens are fixed: apre t,1:ht = (at,1:ht −1 , ct,k ). The algorithm then generates only the remaining suffix at random: apre t,ht +1:H ∼ πpre (· | st , ct,k ). Thus, although the reward function R is deterministic,
Rt = R xit , (at,1:ht −1 , ct,k , apre t,ht +1:H )
is random because the reference-generated suffix is random. For k ∈ {0, 1}, define the expected exponential reward
Mt,k := Ebh +1:H ∼πpre (·|st ,ct,k ) exp t
R(xit , (at,1:ht −1 , ct,k , bht +1:H )) λ
.
The expectation is only over the reference-generated suffix; the prompt, prefix, and candidate token are held fixed. Equivalently, h
i
Mt,k = E eRt /λ Ft , st , ct,1 , ct,0 , Yt = k . To make this expectation explicit, let bht +1:H = (bht +1 , . . . , bH ) denote a possible suffix, whose reference probability is πpre (bht +1:H | st , ct,k ) =
H Y
πpre (bh | xit , (at,1:ht −1 , ct,k , bht +1:h−1 )) .
h=ht +1
By the definition of expectation for a finite distribution, we can compute Mt,k as Mt,k =
X
πpre (bht +1:H | st , ct,k ) × exp
bht +1:H : (at,1:ht −1 ,ct,k ,bht +1:H )∈A(xit )
R(xit , (at,1:ht −1 , ct,k , bht +1:H )) . λ
The feasibility condition ensures that the concatenated sequence is a legal full answer. If ht = H, there is one empty suffix, with probability one.
30
Since R takes values in [0, 1], and the suffix probabilities sum to one, we have that 1 ≤ Mt,k ≤ e1/λ , k ∈ {0, 1}.
(C.3)
We next calculate how acceptance changes the branch distribution. Recall that o
n
It = 1 Ut ≤ e(Rt −1)/λ , Ut ∼ Unif [0, 1], where Ut is independent uniform distribution random variables. Conditioned on Ft , st , ct,1 , ct,0 , Yt , Rt , the acceptance threshold is fixed. Moreover, notice that 0 ≤ Rt ≤ 1 implies e(Rt −1)/λ ∈ [e−1/λ , 1]. By the independence and the uniform density of Ut , we obtain
Pr(It = 1 | Ft , st , ct,1 , ct,0 , Yt , Rt ) = Pr Ut ≤ e(Rt −1)/λ Ft , st , ct,1 , ct,0 , Yt , Rt
=
Z 1 n
o
1 u ≤ e(Rt −1)/λ du = e(Rt −1)/λ .
0
This calculation averages over Ut only; it does not yet average over the reference suffix. Now conditioned only on Ft , st , ct,1 , ct,0 , Yt = k, the selected branch is fixed, but Rt still depends on the random reference suffix. Recall that It is an indicator and we apply the law of iterated expectation to get Pr(It = 1 | Ft , st , ct,1 , ct,0 , Yt = k ) = E[It | Ft , st , ct,1 , ct,0 , Yt = k ]
= E[E[It | Ft , st , ct,1 , ct,0 , Yt , Rt ] | Ft , st , ct,1 , ct,0 , Yt = k ] h
i
= E e(Rt −1)/λ Ft , st , ct,1 , ct,0 , Yt = k . The third uses the acceptance probability computed above. Finally, e(Rt −1)/λ = e−1/λ eRt /λ , and the factor e−1/λ is deterministic. Hence, we have h
i
h
i
E e(Rt −1)/λ Ft , st , ct,1 , ct,0 , Yt = k = e−1/λ E eRt /λ Ft , st , ct,1 , ct,0 , Yt = k = e−1/λ Mt,k . The last equality is the definition of Mt,k . Thus the acceptance probability for branch k is the average of its suffix-specific acceptance probabilities, weighted by the reference suffix law. The branch-selection rule in the algorithm is Pr(Yt = k | Ft , st , ct,1 , ct,0 ) =
πpre (ct,k | st ) . πpre (ct,1 | st ) + πpre (ct,0 | st )
By the property of conditional probability, we have Pr(Yt = k, It = 1 | Ft , st , ct,1 , ct,0 ) = Pr(It = 1|Ft , st , ct,1 , ct,0 , Yt = k ) · Pr(Yt = k|Ft , st , ct,1 , ct,0 )
= e−1/λ
πpre (ct,k | st )Mt,k . πpre (ct,1 | st ) + πpre (ct,0 | st )
All denominators are positive because the reference policy is positive.
31
(C.4)
Summing (C.4) over the two branch indices Yt = 0, 1 and using Mt,k ≥ 1, we obtain πpre (ct,1 | st )Mt,1 + πpre (ct,0 | st )Mt,0 πpre (ct,1 | st ) + πpre (ct,0 | st ) πpre (ct,1 | st ) + πpre (ct,0 | st ) ≥e−1/λ = e−1/λ . πpre (ct,1 | st ) + πpre (ct,0 | st )
Pr(It = 1 | Ft , st , ct,1 , ct,0 ) =e−1/λ
(C.5)
This proves (B.3). By Bayes’ rule, we have that Pr(Yt = 1, It = 1|Ft , st , ct,1 , ct,0 ) Pr(It = 1|Ft , st , ct,1 , ct,0 ) πpre (ct,1 | st )Mt,1 = . πpre (ct,1 | st )Mt,1 + πpre (ct,0 | st )Mt,0
Pr(Yt = 1 | It = 1, Ft , st , ct,1 , ct,0 ) =
(C.6)
It remains to present this distribution in terms of the regularized truth. By Lemma B.1, the optimal policy is πpre (a1:H | x)eR(x,a1:H )/λ πλ⋆ (a1:H | x) = . Zλ ( x ) Its probability of the prefix (at,1:ht −1 , ct,k ) is obtained by summing over all feasible suffixes. The corresponding full-answer events are disjoint, and their union is precisely the prefix event. Hence, by marginalization and Lemma B.1, we have πλ⋆ (at,1:ht −1 , ct,k | xit ) X
=
πλ⋆ ((at,1:ht −1 , ct,k , bht +1:H ) | xit )
bht +1:H : (at,1:ht −1 ,ct,k ,bht +1:H )∈A(xit )
X
=
bht +1:H : (at,1:ht −1 ,ct,k ,bht +1:H )∈A(xit )
R(xit , (at,1:ht −1 , ct,k , bht +1:H )) πpre ((at,1:ht −1 , ct,k , bht +1:H ) | xit ) × exp . Z λ ( x it ) λ
(C.7) We now expand this sum. For every feasible suffix bht +1:H , the autoregressive chain rule separates the full-answer probability into the prefix, the selected token, and the suffix: πpre ((at,1:ht −1 , ct,k , bht +1:H ) | xit )
hY t −1
=
πpre (at,h | xit , at,1:h−1 ) πpre (ct,k | xit , at,1:ht −1 ) ×
H Y
πpre (bh | xit , (at,1:ht −1 , ct,k , bht +1:h−1 )) .
h=ht +1
h=1
The first product is the probability of generating the fixed prefix. The middle factor is the conditional probability of the selected token. The last product is the conditional probability of generating the suffix after that token. More explicitly, using st = (xit , at,1:ht −1 ), we have hY t −1
πpre (at,h | xit , at,1:h−1 ) = πpre (at,1:ht −1 | xit ), πpre (ct,k | xit , at,1:ht −1 ) = πpre (ct,k | st ),
h=1
32
H Y
πpre (bh | xit , (at,1:ht −1 , ct,k , bht +1:h−1 )) = πpre (bht +1:H | st , ct,k ).
h=ht +1
Thus, we have the factorization πpre ((at,1:ht −1 , ct,k , bht +1:H ) | xit ) = πpre (at,1:ht −1 | xit ) πpre (ct,k | st ) πpre (bht +1:H | st , ct,k ). This is a factorization into conditional probabilities, not an independence assumption on the three parts of the answer. At ht = 1 or ht = H, the corresponding empty product is one. Substituting this factorization into (C.7) gives πλ⋆ (at,1:ht −1 , ct,k | xit ) X
=
bht +1:H : (at,1:ht −1 ,ct,k ,bht +1:H )∈A(xit )
πpre (at,1:ht −1 | xit ) πpre (ct,k | st ) × πpre (bht +1:H | st , ct,k ) Zλ ( x i t )
R(xit , (at,1:ht −1 , ct,k , bht +1:H )) × exp . λ
Here the summation variable is only bht +1:H . The quantities πpre (at,1:ht −1 | xit ), πpre (ct,k | st ), Zλ (xit ) do not depend on this suffix. Thus, we factor them out of the sum, and retain exactly the same feasible suffixes to obtain πλ⋆ (at,1:ht −1 , ct,k | xit )
=
πpre (at,1:ht −1 | xit ) πpre (ct,k | st ) Z λ ( x it )
×
X bht +1:H : (at,1:ht −1 ,ct,k ,bht +1:H )∈A(xit )
R(xit , (at,1:ht −1 , ct,k , bht +1:H )) πpre (bht +1:H | st , ct,k ) × exp . λ
The expression in square brackets is precisely the defining suffix sum for Mt,k . Therefore, we get πλ⋆ (at,1:ht −1 , ct,k | xit ) =
πpre (at,1:ht −1 | xit ) πpre (ct,k | st )Mt,k . Z λ ( x it ) π ⋆ (a
,c
|x )
t −1 t,k it . The denominator and Recall the conditional token probability πλ⋆ (ct,k | st ) = λπ⋆ (t,1:h a |x ) λ t,1:ht −1 it the preceding prefix factor are common to both branch indices and are positive. They therefore cancel and we get
πλ⋆ (ct,1 | st ) πpre (ct,1 | st )Mt,1 = . ⋆ ⋆ πλ (ct,1 | st ) + πλ (ct,0 | st ) πpre (ct,1 | st )Mt,1 + πpre (ct,0 | st )Mt,0 Together with (C.6), this proves the first equality in (B.4). Finally, Assumption 3.2 gives πλ⋆ = πwλ⋆ . The common softmax normalizer cancels, so we have ⋆ ⊤
πλ⋆ (ct,1 | st ) 1 e(wλ ) ϕ(st ,ct,1 ) = = σ (zt⊤ wλ⋆ ). = ⋆ )⊤ ϕ(s ,c ⋆ )⊤ ϕ(s ,c ⋆ )⊤ [ϕ(s ,c ⋆ ⋆ ( w ) ( w ) ( − ( w t t,0 t t,1 t t,1 )−ϕ(st ,ct,0 )]) πλ (ct,1 | st ) + πλ (ct,0 | st ) λ e λ +e λ 1+e
33
This establishes teh second equality in (B.4). If ct,1 = ct,0 , the two branch indices have identical reference-completion laws, hence Mt,1 = Mt,0 . Their prior and accepted probabilities are both 1/2, consistent with zt = 0 and σ (0) = 1/2. After EOS, both candidates are null, so the same argument applies. Proof of Lemma B.3. We first recall the random definition vector zt from Algorithm 1. The history Ft contains all randomness before round t, so θt is fixed conditional on Ft . After choosing the source index it and branching time ht , the algorithm draws at,1:H ∼ πtea (· | xit ) and sets st = (xit , at,1:ht −1 ), ct,1 = at,ht is the teacher token, and ct,0 ∼ πstu,θt (· | st ) is a fresh student token at the same state. The vector in the lemma is zt := ϕ(st , ct,1 ) − ϕ(st , ct,0 ) ∈ RD . First, we fix the history Ft and a legal nonterminal state s, so that B (s) is nonempty. Recall that B > 0 is the scalar bound on the parameter norm, whereas B (s) is the set of legal next tokens at s. Both b and c below denote individual tokens, not trajectories. The current student policy is then explicitly
πstu,θt (b | s) =
⊤ϕ exp θ ( s, b ) t stu , X ⊤
b ∈ B (s),
exp θt ϕstu (s, c)
c∈B (s)
b ∈ A \ B (s).
0,
Here θt ∈ Rd is the current student parameter, and ϕstu (s, c) ∈ Rd is the student feature vector of the state-token pair (s, c). Since θt is Ft -measurable, it is fixed after we condition on the history. The model bounds are ∥θt ∥2 ≤ B, ∥ϕstu (s, c)∥2 ≤ 1, ∀c ∈ B (s). Consequently, by the Cauchy-Schwarz inequality, for each legal token c, we have θt⊤ ϕstu (s, c) ≤ ∥θt ∥2 ∥ϕstu (s, c)∥2 ≤ B. Equivalently, −B ≤ θt⊤ ϕstu (s, c) ≤ B. Because u 7→ eu is increasing, we have
e−B ≤ exp θt⊤ ϕstu (s, c) ≤ eB , c ∈ B (s). Now fix any b ∈ B (s). Applying these bounds to the full softmax expression yields
exp θt⊤ ϕstu (s, b)
πstu,θt (b | s) = X
exp θt⊤ ϕstu (s, c)
≥ X
c∈B (s)
c∈B (s)
e−B
≥ X
eB
=
e−B |B (s)|eB
=
e−2B |B (s)|
e−B
exp θt⊤ ϕstu (s, c)
(C.8)
.
c∈B (s)
The first equality is the definition of the student policy. The first inequality replaces its numerator
34
by the lower bound e−B , leaving the positive denominator unchanged. The second inequality increases each denominator term to its upper bound eB ; increasing a positive denominator decreases the fraction. The lower bound applies only to legal tokens; illegal tokens have probability zero. At an absorbing prefix before the terminal horizon, B (s) = {null}, and the softmax formula gives
πstu,θt (null | s) =
exp θt⊤ ϕstu (s, null) exp θt⊤ ϕstu (s, null)
= 1 ≥ e−2B .
Thus (C.8) also holds at these prefixes. To prove the matrix inequality, fix a deterministic vector v ∈ RD and notice v ⊤ zt zt⊤ v = (v ⊤ zt )2 . Conditional expectation is linear, and v is fixed, so we have i
h
i
h
i
h
v ⊤ E zt zt⊤ | Ft v = E v ⊤ zt zt⊤ v | Ft = E (v ⊤ zt )2 | Ft . We now expand this expectation in the order of sampling. The source index it and branching time 1 ht are fresh, independent, uniform draws. Hence, we have Pr(it = i, ht = h | Ft ) = nH . The conditional law of total expectation gives us h
i
E (v ⊤ zt )2 | Ft =
=
H n X X
h
i
Pr(it = i, ht = h | Ft ) E (v ⊤ zt )2 | Ft , it = i, ht = h
i=1 h=1 H n X X
h i 1 E (v ⊤ zt )2 | Ft , it = i, ht = h . nH i=1 h=1
For fixed i, h, let ai,1:h denote a possible realization of the first h teacher tokens, and define si,h := (xi , ai,1:h−1 ). Conditioning on it = i, ht = h, at,1:h = ai,1:h , the definitions of st and ct,1 yield st = si,h and ct,1 = ai,h . The only remaining randomness in zt is the fresh student token. For each b ∈ B (si,h ), its conditional probability is Pr(ct,0 = b | Ft , it = i, ht = h, at,1:h = ai,1:h ) = πstu,θt (b | si,h ). If ct,0 = b, substitution into the recalled definition gives zt = ϕ(si,h , ai,h ) − ϕ(si,h , b). Thus, we have h
E (v ⊤ zt )2 | Ft , it = i, ht = h, at,1:h = ai,1:h
=
X
i
πstu,θt (b | si,h ) v ⊤ [ϕ(si,h , ai,h ) − ϕ(si,h , b)]
2
.
b∈B (si,h )
The frozen teacher’s prefix law is πtea (ai,1:h | xi ) = hk=1 πtea (ai,k | xi , ai,1:k−1 ). The remaining teacher suffix at,h+1:H is absent from zt ; summing its conditional probabilities gives one. Q
For fixed it = i and ht = h, the preceding calculation conditions on the teacher tokens at,1:h = ai,1:h and averages over the student token ct,0 . We now average over all possible teacher token sequences
35
ai,1:h , each weighted by its probability πtea (ai,1:h | xi ). By the law of total expectation, this removes the conditioning on at,1:h and gives E[(v ⊤ zt )2 | Ft , it = i, ht = h]. i
h
v ⊤ E zt zt⊤ | Ft v
n X H 1 X = E nH i=1 h=1 ai,1:h ∼πtea (·|xi )
≥
n X H e−2B X
nH i=1 h=1
= e−2B v ⊤
πstu,θt (b | si,h ) v ⊤ [ϕ(si,h , ai,h ) − ϕ(si,h , b)]
2
Ft
b∈B (si,h )
Eai,1:h ∼πtea (·|xi ) H X
X
1 |B (si,h )|
X
v ⊤ [ϕ(si,h , ai,h ) − ϕ(si,h , b)]
2
Ft
b∈B (si,h )
!
1 Gh v = e−2B v ⊤ Gjoint v. H h=1
The inequality holds by (C.8), and the second equality is by definition (3.7). The final equality P uses Gjoint = H −1 H h = 1 Gh . Here ai,1:h is an integration variable, not an additional rollout: each round still draws only one source index and one teacher answer. Since the calculation holds for every v, it proves the matrix inequality in (B.5). The feature norm bound also gives ∥zt ∥2 = ∥ϕ(st , ct,1 ) − ϕ(st , ct,0 )∥2 ≤ ∥ϕ(st , ct,1 )∥2 + ∥ϕ(st , ct,0 )∥2 ≤ 2. Thus, we have proved (B.5), and we will prove the drift bound (B.6) next. Starting from the definition of the actual stochastic update in Algorithm 1, we have h
i
gt := It zt σ (zt⊤ wt ) − Yt . Recall that It ∈ {0, 1} is the acceptance indicator, Yt ∈ {0, 1} is the selected branch index, and σ (u) = (1 + e−u )−1 . Conditioned on Ft , st , ct,1 , ct,0 , the parameter wt is Ft -measurable, wλ⋆ is fixed, and zt = ϕ(st , ct,1 ) − ϕ(st , ct,0 ) is also determined. Thus, wt , zt , and σ (zt⊤ wt ) are measurable under this conditioning and It and Yt are still randomized. First, by definition of gt , we have n
(wt − wλ⋆ )⊤ gt = (wt − wλ⋆ )⊤ It zt [σ (zt⊤ wt ) − Yt ]
o
i
h
= It (wt − wλ⋆ )⊤ zt [σ (zt⊤ wt ) − Yt ] = zt⊤ (wt − wλ⋆ ) It [σ (zt⊤ wt ) − Yt ]. Taking conditional expectations and pulling out the measurable scalars, we have h
i
h
i
E (wt − wλ⋆ )⊤ gt | Ft , st , ct,1 , ct,0 = zt⊤ (wt − wλ⋆ ) E It [σ (zt⊤ wt ) − Yt ] | Ft , st , ct,1 , ct,0 . We evaluate the remaining expectation by splitting the two possible values of It . When It = 0 the
36
product is zero; when It = 1 it is σ (zt⊤ wt ) − Yt . By the conditional law of total expectation, h
E It [σ (zt⊤ wt ) − Yt ] | Ft , st , ct,1 , ct,0
i h
= Pr(It = 0 | Ft , st , ct,1 , ct,0 ) · 0 + Pr(It = 1 | Ft , st , ct,1 , ct,0 ) × E σ (zt⊤ wt ) − Yt | It = 1, Ft , st , ct,1 , ct,0 h
i
= Pr(It = 1 | Ft , st , ct,1 , ct,0 ) × σ (zt⊤ wt ) − E[Yt | It = 1, Ft , st , ct,1 , ct,0 ] . Because Yt is binary, its conditional expectation is the conditional probability that it equals one. By Lemma B.2, we have that E[Yt | It = 1, Ft , st , ct,1 , ct,0 ]
=0 · Pr(Yt = 0 | It = 1, Ft , st , ct,1 , ct,0 ) + 1 · Pr(Yt = 1 | It = 1, Ft , st , ct,1 , ct,0 ) = Pr(Yt = 1 | It = 1, Ft , st , ct,1 , ct,0 ) = σ (zt⊤ wλ⋆ ). Plugging this back, we obtain h
i
h
i
E (wt − wλ⋆ )⊤ gt | Ft , st , ct,1 , ct,0 = zt⊤ (wt − wλ⋆ ) Pr(It = 1 | Ft , st , ct,1 , ct,0 ) × σ (zt⊤ wt ) − σ (zt⊤ wλ⋆ ) . (C.9) Finally, we pull out the fixed vector zt instead of the scalar inner product, which gives the corresponding vector conditional mean: h
E[gt | Ft , st , ct,1 , ct,0 ] = zt E It [σ (zt⊤ wt ) − Yt ] | Ft , st , ct,1 , ct,0
i
h
i
= Pr(It = 1 | Ft , st , ct,1 , ct,0 ) zt × σ (zt⊤ wt ) − σ (zt⊤ wλ⋆ ) . −u
(C.10)
By direct differentiation, we have σ ′ (u) = (1+ee−u )2 = 2+eu1+e−u > 0. The denominator is even and nondecreasing in |u|, because its derivative on [0, ∞) is eu − e−u ≥ 0. Hence σ ′ (u) ≥ σ ′ (2B ) whenever |u| ≤ 2B. For r ∈ [0, 1], convexity of W places wλ⋆ + r (wt − wλ⋆ ) in W , and therefore we have zt⊤ [wλ⋆ + r (wt − wλ⋆ )] ≤ ∥zt ∥2 ∥wλ⋆ + r (wt − wλ⋆ )∥2 ≤ 2B. Now, we apply the fundamental theorem of calculus along this segment to get σ (zt⊤ wt ) − σ (zt⊤ wλ⋆ ) =
=
Z 1 0
Z 1 0
d ⊤ ⋆ σ zt [wλ + r (wt − wλ⋆ )] dr dr
σ ′ zt⊤ [wλ⋆ + r (wt − wλ⋆ )] zt⊤ (wt − wλ⋆ ) dr
= zt⊤ (wt − wλ⋆ )
37
Z 1 0
σ ′ zt⊤ [wλ⋆ + r (wt − wλ⋆ )] dr.
i
Multiplying the equality by zt⊤ (wt − wλ⋆ ) gives h
i
h
i2 Z 1
h
i2 Z 1
zt⊤ (wt − wλ⋆ ) σ (zt⊤ wt ) − σ (zt⊤ wλ⋆ ) = zt⊤ (wt − wλ⋆ ) ≥ zt⊤ (wt − wλ⋆ )
0
σ ′ zt⊤ [wλ⋆ + r (wt − wλ⋆ )] dr σ ′ (2B ) dr
0
h
= σ ′ (2B ) zt⊤ (wt − wλ⋆ )
i2
≥ 0.
(C.11)
Recall that It = 1{Ut ≤ e(Rt −1)/λ } and that Lemma B.2, (B.3) already establishes Pr(It = 1 | Ft , st , ct,1 , ct,0 ) ≥ e−1/λ . Combining all these together, we obtain i
h
h
E (wt − wλ⋆ )⊤ gt | Ft , st , ct,1 , ct,0 = Pr(It = 1 | Ft , st , ct,1 , ct,0 ) zt⊤ (wt − wλ⋆ ) σ (zt⊤ wt ) − σ (zt⊤ wλ⋆ ) h
≥ Pr(It = 1 | Ft , st , ct,1 , ct,0 ) σ ′ (2B ) zt⊤ (wt − wλ⋆ ) h
≥ e−1/λ σ ′ (2B ) zt⊤ (wt − wλ⋆ )
i2
i
i2
.
The equality was derived by (C.9). The first inequality uses (C.11) and The second one is by (C.5). Now, we apply the tower property once more to obtain h
i
h h
E (wt − wλ⋆ )⊤ gt | Ft = E E (wt − wλ⋆ )⊤ gt | Ft , st , ct,1 , ct,0
h
i2
h
i2
≥ E e−1/λ σ ′ (2B ) zt⊤ (wt − wλ⋆ )
= e−1/λ σ ′ (2B )E zt⊤ (wt − wλ⋆ )
i
Ft
i
Ft
Ft .
i2
h
Notice that zt⊤ (wt − wλ⋆ ) = (wt − wλ⋆ )⊤ zt zt⊤ (wt − wλ⋆ ). Because wt − wλ⋆ is Ft -measurable, by the linearity of conditional expectation, we have E
h
zt⊤ (wt − wλ⋆ )
i2
h
i
Ft = (wt − wλ⋆ )⊤ E zt zt⊤ | Ft (wt − wλ⋆ ).
Using the matrix lower bound (B.5), followed by Gjoint ⪰ µjoint ID , we have h
i
h
i
E (wt − wλ⋆ )⊤ gt Ft ≥ e−1/λ σ ′ (2B )(wt − wλ⋆ )⊤ E zt zt⊤ | Ft (wt − wλ⋆ ) ≥ e−1/λ e−2B σ ′ (2B )(wt − wλ⋆ )⊤ Gjoint (wt − wλ⋆ ) ≥ e−1/λ e−2B σ ′ (2B )µjoint ∥wt − wλ⋆ ∥22
= γ∥wt − wλ⋆ ∥22 . Here ID is the D-dimensional identity matrix. The final equality uses the definition of γ in (4.2). By the Cauchy-Schwarz inequality, we have ∥gt ∥2 = It ∥zt ∥2 |σ (zt⊤ wt ) − Yt | ≤ 2. We prove (B.6). For any u ∈ RD , let p = ProjW (u). Since W is compact and convex, p exists and is unique. For any r ∈ W , the segment p + v (r − p), 0 ≤ v ≤ 1, lies in W . The right derivative at zero of its
38
squared distance to u is nonnegative: 0≤
d = 2⟨p − u, r − p⟩. ∥p + v (r − p) − u∥22 dv v =0+
Thus ⟨u − p, r − p⟩ ≤ 0. Taking r = wλ⋆ and expanding a square gives ∥u − wλ⋆ ∥22 = ∥u − p∥22 + ∥p − wλ⋆ ∥22 + 2⟨u − p, p − wλ⋆ ⟩ ≥ ∥p − wλ⋆ ∥22 , because the cross term and the first squared norm are nonnegative. We use this inequality with u = wt − gt /[γ (t + 2)] and p = wt+1 to get ∥wt+1 − wλ⋆ ∥22 ≤ wt − wλ⋆ −
2 2(wt − wλ⋆ )⊤ gt gt ∥gt ∥22 = ∥wt − wλ⋆ ∥22 − . + 2 γ (t + 2) 2 γ (t + 2) γ (t + 2)2
Conditioning on Ft , we apply (B.5) and (B.6) to get h
E ∥wt+1 − wλ⋆ ∥22
Ft
i
2 4 t 4 ≤ 1− ∥wt − wλ⋆ ∥22 + 2 = ∥wt − wλ⋆ ∥22 + 2 . t+2 γ (t + 2)2 t+2 γ (t + 2)2 (C.12)
We now prove (B.7) by induction. At t = 0 the first coefficient is zero. Taking expectations gives h
i
E ∥w1 − wλ⋆ ∥22 ≤
4 1 ≤ 2. 2 γ 3γ
This is the base case of (B.7). Suppose its assertion holds at some t ≥ 1. By (C.12), we have h
i
E ∥wt+1 − wλ⋆ ∥22 ≤
t 4 4 4(t + 1) 4 + 2 = 2 ≤ 2 . 2 2 2 t + 2 γ (t + 2) γ (t + 2) γ (t + 2) γ (t + 3)
The final inequality follows from (t + 1)(t + 3) = (t + 2)2 − 1 ≤ (t + 2)2 . This proves (B.7). For the increment bound, apply the same projection variational inequality with r = wt ∈ W and the same u, p as above. It implies ∥wt+1 − wt ∥22 ≤ ⟨u − wt , wt+1 − wt ⟩ ≤ ∥u − wt ∥2 ∥wt+1 − wt ∥2 =
∥gt ∥2 ∥wt+1 − wt ∥2 . γ (t + 2)
Thus, we obtain that ∥wt+1 − wt ∥2 ≤ γ (t2+2) for every t ≥ 0. Proof of Lemma B.4. Recall the common feasible sets and the two softmax policies. The set B (s) contains the legal next tokens at a nonterminal prefix s, and A(x) contains the feasible complete answers of length H for prompt x. These sets are finite and do not depend on either parameter. Every nonterminal reachable state has at least one legal token. For a ∈ B (s), πw (a | s) = P
exp(w⊤ ϕ(s, a)) exp(θ⊤ ϕstu (s, a)) P , π ( a | s ) = . stu,θ ⊤ ⊤ b∈B (s) exp(w ϕ(s, b)) b∈B (s) exp(θ ϕstu (s, b))
39
The feature maps are fixed and satisfy ∥ϕ(s, a)∥2 ≤ 1, ∥ϕstu (s, a)∥2 ≤ 1, ∥w∥2 ≤ B, ∥θ∥2 ≤ B. For a fixed feasible answer, sh = (x, a1:h−1 ), the autoregressive definition gives πw (a1:H | x) =
H Y
πw (ah | sh ), πstu,θ (a1:H | x) =
h=1
H Y
πstu,θ (ah | sh ).
h=1
Both probabilities are strictly positive on A(x). Starting from the definition of Zt and the autoregressive product, H πstu,ϑ (a1:H | x) =1 πstu,ϑ (ah | sh ) Zt (ϑ; x, a1:H ) =λ log = λ log QhH πwt (a1:H | x) h=1 πwt (ah | sh )
Q
=λ
H X
[log πstu,ϑ (ah | sh ) − log πwt (ah | sh )]
h=1
=λ
H X
ϑ⊤ ϕstu (sh , ah ) − w⊤ ϕ(sh , ah ) − log
X
t
h=1
ϑ⊤ ϕ
e
stu (sh ,b)
+ log
wt⊤ ϕ(sh ,b)
X
e
b∈B (sh )
b∈B (sh )
We now bound the token log ratios in this exact expression. The calculation is uniform in the parameters, so we carry it out for arbitrary w ∈ W and θ ∈ Θ before substituting w = wt and θ = ϑ. By the Cauchy-Schwarz inequality, for every legal token b, we have that |w⊤ ϕ(s, b)| ≤ ∥w∥2 ∥ϕ(s, b)∥2 ≤ B, |θ⊤ ϕstu (s, b)| ≤ ∥θ∥2 ∥ϕstu (s, b)∥2 ≤ B. Exponentiation preserves these inequalities, so we get e−B ≤ exp(w⊤ ϕ(s, b)) ≤ eB , e−B ≤ exp(θ⊤ ϕstu (s, b)) ≤ eB . Summing the lower and upper bounds over all legal tokens gives |B (s)|e−B ≤
X
exp(w⊤ ϕ(s, b)) ≤ |B (s)|eB , |B (s)|e−B ≤
b∈B (s)
X
exp(θ⊤ ϕstu (s, b)) ≤ |B (s)|eB .
b∈B (s)
The numerators and denominators are positive. Therefore, substituting the preceding bounds yields e−B e−2B eB e2B = , π ( a | s ) ≤ = , w |B (s)|eB |B (s)| |B (s)|e−B |B (s)| e−B e−2B eB e2B πstu,θ (a | s) ≥ = , π = . ( a | s ) ≤ stu,θ |B (s)|eB |B (s)| |B (s)|e−B |B (s)| πw ( a | s ) ≥
Hence, for πstu,θ and πw , we have πstu,θ (a | s) π (a | s) e−2B /|B (s)| e2B /|B (s)| ≥ 2B = e−4B , stu,θ ≤ −2B = e4B . πw (a | s) e /|B (s)| πw ( a | s ) e /|B (s)| We thus obtain −4B ≤ log
πstu,θ (a|s) ≤ 4B. Finally, we have that πw (a|s)
40
(C.13) (C.14)
.
log
H H H Y X X πstu,θ (ah | sh ) πstu,θ (ah | sh ) πstu,θ (ah | sh ) πstu,θ (a1:H | x) log log = log = ≤ ≤ 4BH. πw (a1:H | x) πw (ah | sh ) πw (ah | sh ) πw (ah | sh ) h=1 h=1 h=1
In particular, setting w = wt and θ = ϑ gives |Zt (ϑ; x, a1:H )| = λ log
πstu,ϑ (a1:H | x) ≤ 4λBH. πwt (a1:H | x)
This proves (B.9). Proof of Lemma B.5. Starting from the definition of Cw (θ ), applying Lemma B.4, we have Cw ( θ ) =
≤
≤
m λ X m j =1 m X
λ m j =1
X
ej ) log πstu,θ (a1:H | x
ej ) πstu,θ (a1:H | x ej ) πw (a1:H | x
ej ) log πstu,θ (a1:H | x
ej ) πstu,θ (a1:H | x ej ) πw (a1:H | x
a1:H ∈A(e xj )
X
a1:H ∈A(e xj ) m X X
4λBH m j =1
ej ) = 4λBH. πstu,θ (a1:H | x
a1:H ∈A(e xj )
|
{z
=1
}
For the lower bound, apply the KL non-negativity and we immediately observe that m λ X ej ) ∥ πw (· | x ej )) ≥ 0. Cw (θ ) = KL(πstu,θ (· | x m j =1
Here λ/m > 0, so summing the nonnegative divergences preserves the inequality. Together with the upper bound, this proves (B.10). Now, we bound the derivative of Cw , define the trajectory score Sstu,θ (x, a1:H ) := ∇θ log πstu,θ (a1:H | x). We first bound Sstu,θ . Fix a prompt x and parameter θ, by direct differentiation, we have Sstu,θ (x, a1:H ) = ∇θ log πstu,θ (a1:H | x) = ∇θ log
H Y h=1
41
πstu,θ (ah | sh ) =
H X h=1
∇θ log πstu,θ (ah | sh ).
We evaluate each derivative in this expression, by the linear softmax parametrization, we have
log πstu,θ (ah | sh ) = θ⊤ ϕstu (sh , ah ) − log
X
e
θ⊤ ϕ
stu (sh ,b)
,
b∈B (sh )
P
b∈B (sh ) e
∇θ log πstu,θ (ah | sh ) = ϕstu (sh , ah ) −
θ⊤ ϕstu (sh ,b) ϕ
P
c∈B (sh )
stu (sh , b)
⊤ eθ ϕstu (sh ,c)
⊤
eθ ϕstu (sh ,b) = ϕstu (sh , ah ) − ϕ (s , b) P θ⊤ ϕstu (sh ,c) stu h c∈B (sh ) e b∈B (s ) X
h
X
= ϕstu (sh , ah ) −
πstu,θ (b | sh )ϕstu (sh , b)
b∈B (sh )
=: Dθ,h (x, a1:h ). Substituting the computed token derivative into the score expansion at the start of this step yields Sstu,θ (x, a1:H ) =
H X
∇θ log πstu,θ (ah | sh ) =
h=1
H X
Dθ,h (x, a1:h ).
h=1
The pointwise norm of each increment is bounded by X
∥Dθ,h (x, a1:h )∥2 = ϕstu (sh , ah ) −
πstu,θ (b | sh )ϕstu (sh , b)
b∈B (sh )
≤ ∥ϕstu (sh , ah )∥2 +
2
X
πstu,θ (b | sh )ϕstu (sh , b)
b∈B (sh )
≤ 1+
X
2
πstu,θ (b | sh )∥ϕstu (sh , b)∥2
b∈B (sh )
≤ 1+
X
πstu,θ (b | sh ) = 2.
b∈B (sh )
Consequently, we have ∥Sstu,θ (x, a1:H )∥2 =
H X
≤
Dθ,h (x, a1:h )
h=1
2
H X
∥Dθ,h (x, a1:h )∥2 ≤ 2H.
(C.15)
h=1
To obtain the sharper second-moment bound, we use the conditional centering of the increments.
42
Given x, a1:h−1 , the next token has law ah ∼ πstu,θ (· | sh ). Hence, we have Eah ∼πstu,θ (·|sh ) [Dθ,h (x, a1:h ) | x, a1:h−1 ] X
=
X
πstu,θ (b | sh ) ϕstu (sh , b) −
b∈B (sh )
πstu,θ (c | sh )ϕstu (sh , c)
c∈B (sh )
X
=
πstu,θ (b | sh )ϕstu (sh , b) −
b∈B (sh )
X
=
X
πstu,θ (b | sh )
b∈B (sh )
X
πstu,θ (b | sh )ϕstu (sh , b) −
b∈B (sh )
X
πstu,θ (c | sh )ϕstu (sh , c)
c∈B (sh )
πstu,θ (c | sh )ϕstu (sh , c) = 0.
c∈B (sh )
The second equality uses that the probabilities sum to one. For the conditional second moment, expanding the square using ∥u − v∥22 = ∥u∥22 − 2u⊤ v + ∥v∥22 , we get h
Eah ∼πstu,θ (·|sh ) ∥Dθ,h (x, a1:h )∥22 x, a1:h−1
i 2
=
X
X
πstu,θ (b | sh ) ϕstu (sh , b) −
b∈B (sh )
πstu,θ (c | sh )ϕstu (sh , c)
c∈B (sh )
2
⊤
=
X
πstu,θ (b | sh )∥ϕstu (sh , b)∥22 − 2
πstu,θ (b | sh )ϕstu (sh , b)
X
b∈B (sh )
πstu,θ (b | sh )
πstu,θ (c | sh )ϕstu (sh , c)
2
X
c∈B (sh )
b∈B (sh )
b∈B (sh )
+
X
X
πstu,θ (c | sh )ϕstu (sh , c)
c∈B (sh )
2 2
=
X
πstu,θ (b | sh )∥ϕstu (sh , b)∥22 − 2
b∈B (sh )
X
πstu,θ (c | sh )ϕstu (sh , c)
c∈B (sh )
2
2
X
+
πstu,θ (c | sh )ϕstu (sh , c)
c∈B (sh )
2 2
=
X
πstu,θ (b | sh )∥ϕstu (sh , b)∥22 −
b∈B (sh )
≤
X b∈B (sh )
X
πstu,θ (c | sh )ϕstu (sh , c)
c∈B (sh )
πstu,θ (b | sh )∥ϕstu (sh , b)∥22 ≤
X
2
πstu,θ (b | sh ) = 1.
b∈B (sh )
If h < k, then Dθ,h (x, a1:h ) is already determined by x, a1:k−1 . Marginalizing the suffix after time
43
k and then applying iterated expectation, we get h
Ea1:H ∼πstu,θ (·|x) Dθ,h (x, a1:h )⊤ Dθ,k (x, a1:k )
i
h
i
=Ea1:k−1 ∼πstu,θ (·|x) Eak ∼πstu,θ (·|sk ) Dθ,h (x, a1:h )⊤ Dθ,k (x, a1:k ) x, a1:k−1
=Ea1:k−1 ∼πstu,θ (·|x) Dθ,h (x, a1:h )⊤ Eak ∼πstu,θ (·|sk ) [Dθ,k (x, a1:k ) | x, a1:k−1 ] h
i
=Ea1:k−1 ∼πstu,θ (·|x) Dθ,h (x, a1:h )⊤ 0 = 0. Pulling the earlier increment outside the inner expectation is valid precisely because the earlier increment is measurable given that prefix. By the linearity of conditional means, we further have Ea1:H ∼πstu,θ (·|x) [Sstu,θ (x, a1:H )] =
H X
Ea1:h−1 ∼πstu,θ (·|x) Eah ∼πstu,θ (·|sh ) [Dθ,h (x, a1:h ) | x, a1:h−1 ]
h=1
=
H X
Ea1:h−1 ∼πstu,θ (·|x) [0] = 0.
h=1
When h = 1, the prefix is empty and its expectation has just one possible value. Expanding the square of the sum of the increments, including all cross terms, gives h
Ea1:H ∼πstu,θ (·|x) ∥Sstu,θ (x, a1:H )∥22
=Ea1:H ∼πstu,θ (·|x)
H X
2
=
Dθ,h (x, a1:h )
h=1 H X
i
h
2
i
Ea1:H ∼πstu,θ (·|x) ∥Dθ,h (x, a1:h )∥22 + 2
h=1
=
X
h
Ea1:H ∼πstu,θ (·|x) Dθ,h (x, a1:h )⊤ Dθ,k (x, a1:k )
1≤h<k≤H |
{z
}
=0
H X
i
h
i
Ea1:h−1 ∼πstu,θ (·|x) Eah ∼πstu,θ (·|sh ) ∥Dθ,h (x, a1:h )∥22 x, a1:h−1
h=1
≤
H X
Ea1:h−1 ∼πstu,θ (·|x) [1] = H.
h=1
We have therefore proved the two score identities h
Ea1:H ∼πstu,θ (·|x) [Sstu,θ (x, a1:H )] = 0, Ea1:H ∼πstu,θ (·|x) ∥Sstu,θ (x, a1:H )∥22
44
i
≤ H.
(C.16)
Now, we can bound ∇θ Cw (θ ). Fix w while differentiating with respect to θ.
m λ X ∇ θ Cw ( θ ) = ∇ θ m j =1
a1:H ∈A(e xj )
ej ) πstu,θ (a1:H | x ej ) log πstu,θ (a1:H | x ej ) πw (a1:H | x
m X
λ = m j =1
X
X a1:H ∈A(e xj )
ej ) πstu,θ (a1:H | x . ej ) log ∇θ πstu,θ (a1:H | x e πw (a1:H | xj )
We now expand the derivative of each summand. First, by direct algebra, for any x, we have ∇θ πstu,θ (a1:H | x) = ∇θ exp[log πstu,θ (a1:H | x)] = πstu,θ (a1:H | x)∇θ log πstu,θ (a1:H | x)
= πstu,θ (a1:H | x)Sstu,θ (x, a1:H ). Moreover, πw does not depend on θ, and hence we have ∇θ log
πstu,θ (a1:H | x) = ∇θ log πstu,θ (a1:H | x) − ∇θ log πw (a1:H | x) = Sstu,θ (x, a1:H ) − 0. πw (a1:H | x)
Therefore, by the product rule of differentiation, we have πstu,θ (a1:H | x) πw (a1:H | x) π (a1:H | x) π (a1:H | x) = [∇θ πstu,θ (a1:H | x)] log stu,θ + πstu,θ (a1:H | x)∇θ log stu,θ πw (a1:H | x) πw (a1:H | x) π (a1:H | x) =πstu,θ (a1:H | x)Sstu,θ (x, a1:H ) log stu,θ + πstu,θ (a1:H | x)Sstu,θ (x, a1:H ). πw (a1:H | x)
∇θ πstu,θ (a1:H | x) log
For the second term, its sum is the score expectation already evaluated in (C.16): πstu,θ (a1:H | x)Sstu,θ (x, a1:H ) = Ea1:H ∼πstu,θ (·|x) [Sstu,θ (x, a1:H )] = 0.
X a1:H ∈A(x)
The first equality is the definition of expectation and the second is the first identity in (C.16), applied to this x and θ. Substituting both product-rule terms into the definition of Cw , and then using this zero sum, we obtain ∇ θ Cw ( θ ) =
+
m λ X m j =1
X a1:H ∈A(e xj )
m λ X
m j =1
ej )Sstu,θ (x ej , a1:H ) × log πstu,θ (a1:H | x
X
ej ) πstu,θ (a1:H | x ej ) πw (a1:H | x
ej )Sstu,θ (x ej , a1:H ) πstu,θ (a1:H | x
a1:H ∈A(e xj )
|
{z
=0
}
m ej ) πstu,θ (a1:H | x λ X . ej , a1:H ) × log = Ea1:H ∼πstu,θ (·|exj ) Sstu,θ (x ej ) m j =1 πw (a1:H | x
In the last line, we rewrite the finite weighted sum as an expectation. The quantity we want to control has the exact representation
45
(C.17)
|Cw (θ ) − Cw (ϑ)| =
Z 1 0
Z 1
d Cw ϑ + u(θ − ϑ) du = du
∇ θ Cw ϑ + u ( θ − ϑ )
⊤
(θ − ϑ) du .
0
(C.18) These identities follow from the fundamental theorem of calculus. We bound the gradient appearing here first and then return to this expression. By the triangle inequality, the pointwise log-ratio bound in Lemma B.4, we have ∥∇θ Cw (θ )∥2 =
m λ X m j =1
X
ej )Sstu,θ (x ej , a1:H ) × log πstu,θ (a1:H | x
a1:H ∈A(e xj )
ej ) πstu,θ (a1:H | x e πw (a1:H | xj )
2
m X
≤
ej ) πstu,θ (a1:H | x λ ej , a1:H )∥2 × log Ea1:H ∼πstu,θ (·|exj ) ∥Sstu,θ (x ej ) m j =1 πw (a1:H | x
≤
m 4λBH X E [∥Sstu,θ (xej , a1:H )∥2 ] m j =1 a1:H ∼πstu,θ (·|exj )
≤
m h i1/2 4λBH X ej , a1:H )∥22 Ea1:H ∼πstu,θ (·|exj ) ∥Sstu,θ (x m j =1
≤
m √ 4λBH X H = 4λBH 3/2 . m j =1
In the second inequality, we apply Lemma B.4 and in the last inequality, we apply (C.16). For u ∈ [0, 1], the segment between ϑ and θ remains inside the student ball because ∥ϑ + u(θ − ϑ)∥2 = ∥(1 − u)ϑ + uθ∥2 ≤ (1 − u)∥ϑ∥2 + u∥θ∥2 ≤ (1 − u)B + uB = B. Thus, the gradient bound applies at every point of the segment in (C.18). Return to that representation and we get |Cw (θ ) − Cw (ϑ)| ≤
Z 1
∇ θ Cw ϑ + u ( θ − ϑ )
⊤
(θ − ϑ) du
0
≤ ∥θ − ϑ∥2
Z 1
∇ θ Cw ϑ + u ( θ − ϑ )
0
≤ ∥θ − ϑ∥2
Z 1
2
du
4λBH 3/2 du = 4λBH 3/2 ∥θ − ϑ∥2 .
0
This proves (B.11). To prove (B.12), we start with the difference that must be bounded and expand
46
two costs, we have Cw ( θ ) − C v ( θ ) =
=
=
m λ X m j =1
X a1:H ∈A(e xj )
m λ X
m j =1 m λ X m j =1
ej ) × [log πstu,θ (a1:H | x
ej ) ej ) πstu,θ (a1:H | x πstu,θ (a1:H | x − log ] ej ) ej ) πw (a1:H | x πv (a1:H | x
X
ej ) × − log πw (a1:H | x ej ) + log πv (a1:H | x ej ) πstu,θ (a1:H | x
a1:H ∈A(e xj )
X
ej ) × [log πv (a1:H | x ej ) − log πw (a1:H | x ej )] . πstu,θ (a1:H | x
a1:H ∈A(e xj )
(C.19) Thus, it suffices to control the difference of calibration log probabilities. We derive that control from the calibration softmax itself:
log πw (a | s) = w⊤ ϕ(s, a) − log
X
w⊤ ϕ(s,b)
e
,
b∈B (s)
P
b∈B (s) e
∇w log πw (a | s) =ϕ(s, a) −
w⊤ ϕ(s,b) ϕ(s, b)
w⊤ ϕ(s,c) c∈B (s) e
P X
=ϕ(s, a) −
⊤
ew ϕ(s,b) = ϕ(s, a) − ϕ(s, b) P w⊤ ϕ(s,c) c∈B (s) e b∈B (s) X
πw (b | s)ϕ(s, b).
b∈B (s)
We bound this derivative by the following: ∥∇w log πw (a | s)∥2 = ϕ(s, a) −
X
πw (b | s)ϕ(s, b)
b∈B (s)
≤ 1+
X
≤ ∥ϕ(s, a)∥2 +
πw (b | s)ϕ(s, b)
b∈B (s)
2
πw (b | s)∥ϕ(s, b)∥2 ≤ 1 +
X
X
2
πw (b | s) = 2.
b∈B (s)
b∈B (s)
Since W is convex, v + u(w − v ) ∈ W for u ∈ [0, 1]. The same one-dimensional integration argument, now displayed for the token log probability, gives log πw (a | s) − log πv (a | s) =
Z 1 0
=
Z 1 0
d log πv +u(w−v ) (a | s) du du ∇w log πw (a | s)|⊤ w =v +u(w−v ) (w − v ) du.
Consequently, we have | log πw (a | s) − log πv (a | s)| ≤
Z 1 0
≤
Z 1
∇w log πw (a | s)|w=v +u(w−v ) 2∥w − v∥2 du = 2∥w − v∥2 .
0
47
2
∥w − v∥2 du
Thus, we have H h X
log πw (ah | sh ) − log πv (ah | sh )
| log πw (a1:H | x) − log πv (a1:H | x)| =
i
h=1
≤ ≤
H X h=1 H X
| log πw (ah | sh ) − log πv (ah | sh )| 2∥w − v∥2 = 2H∥w − v∥2 .
h=1
Now, we return to the cost-difference representation (C.19). Plugging the bound derived above back, we obtain, |Cw (θ ) − Cv (θ )| ≤ ≤
=
m λ X m j =1
X
ej ) × | log πv (a1:H | x ej ) − log πw (a1:H | x ej ) | πstu,θ (a1:H | x
a1:H ∈A(e xj ) m 2λH∥w − v∥2 X
m
X
ej ) πstu,θ (a1:H | x
j =1 a1:H ∈A(e xj )
m 2λH∥w − v∥2 X 1 = 2λH∥w − v∥2 . m j =1
The right-hand side is independent of θ. Taking the supremum over θ ∈ Θ proves (B.12). Proof of Lemma B.7. We first verify analyticity, which will also be used in the two Lojasiewicz arguments below. For each fixed legal state s and token a ∈ B (s), the softmax expressions are ⊤
eθ ϕstu (s,a) , πstu,θ (a | s) = P θ⊤ ϕstu (s,b) b∈B (s) e
log πstu,θ (a | s) = θ⊤ ϕstu (s, a) − log
X
⊤
eθ ϕstu (s,b) .
b∈B (s)
The exponentials of these linear functions are analytic, and their finite sum is strictly positive on Rd . Since reciprocals and logarithms are analytic on (0, ∞), both displayed expressions are analytic. For each fixed feasible answer, the chain rule for trajectory probabilities gives πstu,θ (a1:H | x) = log πstu,θ (a1:H | x) =
H Y h=1 H X
πstu,θ (ah | x, a1:h−1 ), log πstu,θ (ah | x, a1:h−1 ).
h=1
Finite products and sums preserve analyticity. For any fixed positive answer laws Qj , expanding the average KL gives m m 1 X 1 X ej ) ∥ Qj ) = KL(πstu,θ (· | x m j =1 m j =1
X
ej ) × [log πstu,θ (a1:H | x ej ) − log Qj (a1:H )] . πstu,θ (a1:H | x
a1:H ∈A(e xj )
48
Each log Qj (a1:H ) is a finite constant, and the sum has finitely many analytic summands. Thus this ej ) and multiplying by λ proves the assertion for Cw ; function is analytic. Choosing Qj = πw (· | x † ej ) proves it for Kλ,m . In particular, F = Cw⋆ and ∆λ,m = F − F (θλ,m choosing Qj = πstu,θ† (· | x ) λ λ,m
are analytic. We next bound the Hessian of Cw . For this calculation, define the unscaled trajectory log ratio ℓw (θ; x, a1:H ) := log
πstu,θ (a1:H | x) . πw (a1:H | x)
By Lemma B.4, |ℓw | ≤ 4BH on Θ. Recall from (C.17) that ∇ θ Cw ( θ ) =
m λ X m j =1
X
ej )ℓw (θ; x ej , a1:H )Sstu,θ (x ej , a1:H ). πstu,θ (a1:H | x
a1:H ∈A(e xj )
We take the derivative again and use the product rule to compute the Hessian matrix ∇2θ Cw (θ ) =
m λ X ej , a1:H ) + 1 Sstu,θ (x ej , a1:H )Sstu,θ (x ej , a1:H )⊤ Ea1:H ∼πstu,θ (·|exj ) ℓw (θ; x m j =1
+ ℓw (θ; xej , a1:H )∇θ Sstu,θ (xej , a1:H ) .
(C.20)
We already know from (C.16) that the expected squared score norm is at most H. To bound its derivative, we differentiate the token-score formula in the proof of Lemma B.5: ∇θ Sstu,θ (x, a1:H ) = −
H X
X
h=1
πstu,θ (b | sh )ϕstu (sh , b)ϕstu (sh , b)⊤
b∈B (sh )
−
⊤
X
πstu,θ (b | sh )ϕstu (sh , b)
X
πstu,θ (b | sh )ϕstu (sh , b) .
b∈B (sh )
b∈B (sh )
This is a covariance matrix and is positive semidefinite. For any unit vector v ∈ Rd , we have 0 ≤ −v ⊤ ∇θ Sstu,θ (x, a1:H )v
=
H X
H X
2
πstu,θ (b | sh ) v ⊤ ϕstu (sh , b)
X
h=1
≤
2
−
b∈B (sh )
X
X
πstu,θ (b | sh )v ⊤ ϕstu (sh , b)
b∈B (sh )
πstu,θ (b | sh )∥v∥22 ∥ϕstu (sh , b)∥22 ≤
h=1 b∈B (sh )
H X
1 = H.
h=1
The first upper bound discards a nonnegative square and applies Cauchy-Schwarz; the last uses the feature norm bound and that the token probabilities sum to one. Thus ∥∇θ Sstu,θ (x, a1:H )∥op ≤ H.
49
Taking operator norms in (C.20) and using ∥SS ⊤ ∥op = ∥S∥22 now yields ∥∇2θ Cw (θ )∥op ≤ ≤
m i h λ X ej , a1:H )∥22 + 4BH 2 Ea1:H ∼πstu,θ (·|exj ) (4BH + 1)∥Sstu,θ (x m j =1 m h i λ X (4BH + 1)H + 4BH 2 = λH (1 + 8BH ) = Lst . m j =1
This proves (B.15). Since Θ is convex, the segment joining ϑ and θ lies in Θ. The fundamental theorem of calculus then gives ∥∇θ Cw (θ ) − ∇θ Cw (ϑ)∥2 =
Z 1 0
∇2θ Cw (ϑ + u(θ − ϑ))(θ − ϑ) du
≤
Z 1
Lst ∥θ − ϑ∥2 du = Lst ∥θ − ϑ∥2 ,
0
2
which proves (B.16). Finally, subtract the two representations in (C.17). The student law and its score agree in both terms, so cancellation of their log probabilities gives m ej ) λ X πv (a1:H | x ej , a1:H ) log ∇ θ Cw ( θ ) − ∇ θ Cv ( θ ) = Ea1:H ∼πstu,θ (·|exj ) Sstu,θ (x . ej ) m j =1 πw (a1:H | x
"
#
The proof of (B.12) already established the pointwise bound | log(πv (a1:H | x)/πw (a1:H | x))| ≤ 2H∥w − v∥2 . Applying that bound and Cauchy-Schwarz to the score expectation, and then using (C.16), we obtain ∥∇θ Cw (θ ) − ∇θ Cv (θ )∥2 ≤ ≤
m 2λH∥w − v∥2 X ej , a1:H )∥2 ] Ea1:H ∼πstu,θ (·|exj ) [∥Sstu,θ (x m j =1 m h i1/2 2λH∥w − v∥2 X ej , a1:H )∥22 Ea1:H ∼πstu,θ (·|exj ) ∥Sstu,θ (x m j =1
≤ 2λH 3/2 ∥w − v∥2 . This proves (B.17). Proof of Lemma B.6. We first compute the conditional mean of each summand, then bound the variance of their average. Conditional on Gt+1 , θt and wt+1 are fixed. For every ℓ, Algorithm 1 samples jtg+1,ℓ uniformly and then samples an answer from the current student at that question. ej ), we have Thus, for a1:H ∈ A(x
Pr jtg+1,ℓ = j, agt+1,ℓ,1:H = a1:H Gt+1 = Pr jtg+1,ℓ = j Gt+1 × Pr agt+1,ℓ,1:H = a1:H Gt+1 , jtg+1,ℓ = j
=
1 ej ) . πstu,θt (a1:H | x m
50
(C.21)
Starting from the definition of gbtstu +1,ℓ , we obtain h
i
h
ej g E gbtstu +1,ℓ Gt+1 = E Sstu,θt (x
t+1,ℓ
m 1 X = m j =1
X
λ m j =1
t+1,ℓ
, agt+1,ℓ,1:H ) Gt+1
i
ej )Sstu,θt (x ej , a1:H )Zt+1 (θt ; x ej , a1:H ) πstu,θt (a1:H | x
a1:H ∈A(e xj )
m X
=
ej g , agt+1,ℓ,1:H )Zt+1 (θt ; x
X
ej )Sstu,θt (x ej , a1:H ) log πstu,θt (a1:H | x
a1:H ∈A(e xj )
ej ) πstu,θt (a1:H | x ej ) πwt+1 (a1:H | x
= ∇θ Cwt+1 (θt ) = ∇θ Ct+1 (θt ). The second equality uses (C.21), the third substitutes the definition of Zt+1 , and the fourth applies the cost-gradient identity (C.17). By linearity of conditional expectation, we have, h
i
E gbtstu + 1 Gt + 1 =
bt + 1 h i 1 X bt+1 E gbtstu G = ∇ θ Ct + 1 ( θ t ) = ∇ θ Ct + 1 ( θ t ) . t + 1 +1,ℓ bt+1 ℓ=1 bt+1
To control the variance, we first bound the second moment. Its definition and (C.21) give h
i
2 E ∥gbtstu +1,ℓ ∥2 Gt+1 =
≤
m 1 X m j =1
X
a1:H ∈A(e xj ) m 2 2 2 16λ B H X
m
ej , a1:H )∥22 [Zt+1 (θt ; x ej , a1:H )]2 ej )∥Sstu,θt (x πstu,θt (a1:H | x X
ej , a1:H )∥22 ej )∥Sstu,θt (x πstu,θt (a1:H | x
j =1 a1:H ∈A(e xj )
=
m h i 16λ2 B 2 H 2 X ej , a1:H )∥22 Ea1:H ∼πstu,θ (·|exj ) ∥Sstu,θt (x t m j =1
≤
m 16λ2 B 2 H 2 X H = 16λ2 B 2 H 3 . m j =1
The first inequality uses |Zt+1 | ≤ 4λBH from Lemma B.4. The last inequality applies the score second-moment bound (C.16) at each target question. For this proof, we set et+1,ℓ := gbtstu +1,ℓ − ∇θ Ct+1 (θt ). The identity above yields E[et+1,ℓ | Gt+1 ] = 0. Expanding the square yields D
E
2 2 btstu E[∥et+1,ℓ ∥22 | Gt+1 ] = E[∥gbtstu +1,ℓ ∥2 | Gt+1 ] − 2 E[g +1,ℓ | Gt+1 ], ∇θ Ct+1 (θt ) + ∥∇θ Ct+1 (θt )∥2 2 2 2 2 3 = E[∥gbtstu +1,ℓ ∥2 | Gt+1 ] − ∥∇θ Ct+1 (θt )∥2 ≤ 16λ B H .
For ℓ ̸= ℓ′ , the two prompt-answer pairs are independent conditional on Gt+1 , and therefore E[⟨et+1,ℓ , et+1,ℓ′ ⟩ | Gt+1 ] = E[et+1,ℓ | Gt+1 ], E[et+1,ℓ′ | Gt+1 ] = 0.
51
Using the definition of the averaged estimator, we conclude that
E gbtstu + 1 − ∇ θ Ct + 1 ( θ t )
2
Gt+1 = E
2
1
bX t+1
bt+1 ℓ=1
2
Gt+1
et+1,ℓ
2
bX t+1
=
1 2 E[∥et+1,ℓ ∥22 | Gt+1 ] + 2 E[⟨et+1,ℓ , et+1,ℓ′ ⟩ | Gt+1 ] 2 bt+1 ℓ=1 bt+1 ℓ<ℓ′
≤
bt+1 (16λ2 B 2 H 3 ) 16λ2 B 2 H 3 . = bt+1 b2t+1
X
This proves (B.14). Finally, θt is Ft -measurable, so σ (wt+1 , θt ) ⊆ Gt+1 . The tower property gives the remaining conditional identity: h
i
h
i
btstu E gbtstu +1 wt+1 , θt = E E[g +1 | Gt+1 ] wt+1 , θt = E[∇θ Cwt+1 (θt ) | wt+1 , θt ] = ∇θ Ct+1 (θt ).
The last equality holds because the displayed gradient is a function of wt+1 and θt . This completes the proof. Proof of Lemma B.8. Conditional on Ht , the candidate policies and wt are fixed. Since we sample ejt,k,ℓ ), we have jt,k,ℓ uniformly from {1, . . . , m} and then generate aval t,k,ℓ,1:H from πstu,ϑt,k (· | x h
ejt,k,ℓ , aval E Zt (ϑt,k ; x t,k,ℓ,1:H ) | Ht
=
m 1 X m j =1
X
i
ej )λ log πstu,ϑt,k (a1:H | x
a1:H ∈A(e xj )
ej ) πstu,ϑt,k (a1:H | x = Ct (ϑt,k ). ej ) πwt (a1:H | x
2 2 2 ejt,k,ℓ , aval By Lemma B.4, the squared summand of Zt (ϑt,k ; x t,k,ℓ,1:H ), ℓ = 1, · · · , qt is at most 16λ B H . Since the samples within a batch are independent and centered after subtracting their mean, the cross terms vanish when expanding the squared sample-mean error and we have
1 E[(Cbt,k − Ct (ϑt,k ))2 | Ht ] =
qt X
qt2 ℓ=1
ejt,k,ℓ , aval Var Zt (ϑt,k ; x t,k,ℓ,1:H ) | Ht ≤
16λ2 B 2 H 2 . qt
Then, we apply the Cauchy-Schwarz inequality and get
E[|Cbt,k − Ct (ϑt,k )| | Ht ] ≤ E[(Cbt,k − Ct (ϑt,k ))2 | Ht ]
1/2
4λBH ≤ √ . qt
The maximum of two nonnegative numbers is no larger than their sum. Consequently, we have E[ξt | Ht ] ≤ 2
2 X
E[|Cbt,k − Ct (ϑt,k )| | Ht ] ≤
k =1
16λBH 16λBH = . √ qt t+1
Choose an index kt⋆ that minimizes the two true costs, only for this proof. The definition of ξt and
52
empirical minimization implies that Ct (θt ) = Ct (ϑt,bk ) ≤ Cbt,bk + ξt /2 ≤ Cbt,kt⋆ + ξt /2 ≤ Ct (ϑt,kt⋆ ) + ξt . t
t
(C.22)
The first and last inequalities bound the evaluation errors; the middle inequality is the actual two-way selection rule. This proves (B.18) and we finish the proof. † Proof of Lemma B.9. Recall that F = Cwλ⋆ and ∆λ,m (θ ) = F (θ ) − F (θλ,m ). By optimality, the student Lipschitz bound in (B.11), and the diameter 2B of Θ, † 0 ≤ ∆λ,m (θ ) ≤ 4λBH 3/2 ∥θ − θλ,m ∥2 ≤ 8λB 2 H 3/2 = Mopt .
(C.23)
We first show that the exact projected step gives a quadratic improvement when this gap is small, and then use uniform exploration for the remaining values of the gap. To apply Lemma A.2, we define the extended-real function and the normal cone (
Ψ (θ ) : =
F (θ ), +∞,
θ ∈ Θ, θ∈ / Θ,
NΘ (θ ) := {v ∈ Rd : v ⊤ (ϑ − θ ) ≤ 0 for every ϑ ∈ Θ}, θ ∈ Θ. For this smooth function on a closed convex set, We claim that ∂Ψ(θ ) = ∇θ F (θ ) + NΘ (θ ), θ ∈ Θ.
(C.24)
For θ, ϑ ∈ Θ, we first have F (ϑ) − F (θ ) = ∇θ F (θ )⊤ (ϑ − θ ) + o(∥ϑ − θ∥2 ). Hence, by the definition of the Fréchet subdifferential, we have b (θ ) ⇐⇒ v ∈ ∂Ψ
F (ϑ) − F (θ ) − v ⊤ (ϑ − θ ) ( ∇θ F ( θ ) − v ) ⊤ ( ϑ − θ ) ≥ 0 ⇐⇒ lim inf ≥ 0. ϑ→θ, ϑ∈Θ ϑ→θ, ϑ∈Θ ∥ϑ − θ∥2 ∥ϑ − θ∥2 lim inf ϑ̸=θ
ϑ̸=θ
On one hand, by the definition of NΘ (θ ), we know that b (θ ). v − ∇θ F (θ ) ∈ NΘ (θ ) =⇒ (∇θ F (θ ) − v )⊤ (ϑ − θ ) ≥ 0 ∀ϑ ∈ Θ =⇒ v ∈ ∂Ψ b (θ ). ∀ζ ∈ Θ \ {θ}, by the convexity of Θ, we have that Conversely, let v ∈ ∂Ψ
ϑu = θ + u(ζ − θ ) ∈ Θ, u ∈ (0, 1]. Therefore, 0 ≤ lim inf u↓0
( ∇ θ F ( θ ) − v ) ⊤ ( ϑu − θ ) u ( ∇θ F ( θ ) − v ) ⊤ ( ζ − θ ) ( ∇θ F ( θ ) − v ) ⊤ ( ζ − θ ) = lim inf = . u↓0 ∥ϑu − θ∥2 u∥ζ − θ∥2 ∥ζ − θ∥2
Consequently, we obtain (v − ∇θ F (θ ))⊤ (ζ − θ ) ≤ 0 ∀ζ ∈ Θ. Thus, v − ∇θ F (θ ) ∈ NΘ (θ ). b (θ ) = ∇ F (θ ) + N (θ ). By the definition of the limiting Combining both inclusions yields ∂Ψ θ Θ
53
subdifferential, we know that
v ∈ ∂Ψ(θ ) ⇐⇒ ∃ (θn , vn )n≥1 :
θn ∈ Θ, θn → θ,
Ψ(θ ) → Ψ(θ ), vn → v,
n v ∈ ∂Ψ b ( θn ) . n
For such a sequence and every ζ ∈ Θ, we have (vn − ∇θ F (θn ))⊤ (ζ − θn ) ≤ 0. Continuity of ∇θ F yields that
(v − ∇θ F (θ ))⊤ (ζ − θ ) = lim (vn − ∇θ F (θn ))⊤ (ζ − θn ) ≤ 0. n→∞
Thus, we obtain ∂Ψ(θ ) ⊆ ∇θ F (θ ) + NΘ (θ ). For the reverse inclusion, the constant sequences θn = θ and vn = v give b (θ ) =⇒ v ∈ ∂Ψ(θ ). v ∈ ∇θ F (θ ) + NΘ (θ ) = ∂Ψ
Therefore, we have proved the claim. Lemma B.7 shows that F is analytic. Consequently, the finite graph of Ψ is {(θ, u) ∈ Rd × R : ∥θ∥22 ≤ B 2 , u − F (θ ) = 0}, which is semianalytic and hence subanalytic. Its domain is the closed ball Θ, and Ψ is continuous † on this domain. Moreover, optimality of θλ,m on every segment in Θ implies that † † ∇θ F (θλ,m )⊤ (ϑ − θλ,m ) ≥ 0, ϑ ∈ Θ. † † † ). All hypotheses of Lemma A.2 ), and (C.24) gives 0 ∈ ∂Ψ(θλ,m ) ∈ NΘ (θλ,m Thus −∇θ F (θλ,m † therefore hold. There are an open neighborhood U of θλ,m , cloc > 0, and ρloc ∈ [0, 1) such that
dist(0, ∂Ψ(θ )) ≥ cloc [∆λ,m (θ )]ρloc , θ ∈ U ∩ Θ, ∆λ,m (θ ) > 0.
(C.25)
Uniqueness of the minimizer now lets us choose 0 < δloc ≤ min{1, Mopt } such that {θ ∈ Θ : ∆λ,m (θ ) ≤ δloc } ⊂ U . To justify this choice, if Θ \ U is nonempty, its compactness and continuity of ∆λ,m give an attained † minimum there. This minimum is strictly positive because the only zero is θλ,m ∈ U . Choose δloc smaller than that minimum. If the complement is empty, any positive value satisfying the displayed upper bound suffices. For 0 < ∆λ,m (θ ) ≤ δloc ≤ 1, the fact that ρloc < 1 yields
[∆λ,m (θ )]ρloc = ∆λ,m (θ )[∆λ,m (θ )]ρloc −1 ≥ ∆λ,m (θ ). Hence (C.25) implies the weaker but sufficient linear bound dist(0, ∂Ψ(θ )) ≥ cloc ∆λ,m (θ ), 0 < ∆λ,m (θ ) ≤ δloc .
(C.26)
We apply this bound to the exact projected step y (θ ) defined in the lemma. The first-order 54
condition for Euclidean projection is
θ − αstu ∇θ F (θ ) − y (θ )
⊤
(ϑ − y (θ )) ≤ 0 ∀ϑ ∈ Θ.
Setting ϑ = θ and rearranging gives ∇θ F ( θ ) ⊤ ( y ( θ ) − θ ) ≤ −
∥y (θ ) − θ∥22 . αstu
The gradient Lipschitz bound in (B.16), integrated along the segment from θ to y (θ ), gives F (y (θ )) − F (θ ) = ∇θ F (θ )⊤ (y (θ ) − θ ) +
Z 1h
∇θ F (θ + u(y (θ ) − θ )) − ∇θ F (θ )
i⊤
(y (θ ) − θ ) du
0 1 ∥y (θ ) − θ∥22 uLst ∥y (θ ) − θ∥22 du ≤− + αstu 0 1 Lst ∥y (θ ) − θ∥22 =− ∥y (θ ) − θ∥22 ≤ − − . αstu 2 2αstu
Z
The last inequality uses αstu = 1/(2Lst ). In particular, F (y (θ )) ≤ F (θ ). The projection condition also implies θ − y (θ ) − ∇θ F (θ ) ∈ NΘ (y (θ )). αstu Using (C.24) at y (θ ), and then (B.16), we obtain 1 θ − y (θ ) ≤ Lst + ∥y (θ ) − θ∥2 . dist(0, ∂Ψ(y (θ ))) ≤ ∇θ F (y (θ )) − ∇θ F (θ ) + αstu αstu 2
Suppose now that 0 < ∆λ,m (θ ) ≤ δloc . If ∆λ,m (y (θ )) ≤ ∆λ,m (θ )/2, then (C.23) gives F (θ ) − F (y (θ )) = ∆λ,m (θ ) − ∆λ,m (y (θ )) ≥
∆λ,m (θ ) [∆λ,m (θ )]2 ≥ . 2 2Mopt
Otherwise, 0 < ∆λ,m (θ )/2 < ∆λ,m (y (θ )) ≤ δloc , so (C.26) applies at y (θ ). Combining the preceding two estimates with that inequality gives ∥y (θ ) − θ∥22 dist(0, ∂Ψ(y (θ )))2 ≥ 2αstu 2αstu (Lst + 1/αstu )2 c2loc [∆λ,m (y (θ ))]2 c2loc ≥ ≥ [∆λ,m (θ )]2 . 2αstu (Lst + 1/αstu )2 8αstu (Lst + 1/αstu )2
F (θ ) − F (y (θ )) ≥
The same quadratic lower bound, with the smaller of the two coefficients, therefore holds throughout the small sublevel set. When the gap is zero, the required bound is immediate. It remains to control ∆λ,m (θ ) > δloc . We establish the needed exploration probability for a calibrated objective. Fix w ∈ W and 0 < u ≤ Mopt . By continuity and compactness, we can choose ϑ⋆w ∈ argmin Cw (ϑ). ϑ∈Θ
Set τ = u/Mopt ∈ (0, 1] and consider {(1 − τ )ϑ⋆w + τ ζ : ζ ∈ Θ}. By convexity, we know that this
55
set lies in Θ. Each of its points ϑ has distance at most 2Bτ from ϑ⋆w , so (B.11) gives Cw (ϑ) − min Cw (ζ ) ≤ (4λBH 3/2 )(2Bτ ) = Mopt τ = u. ζ∈Θ
The affine map defining the set scales d-dimensional volume by τ d . Since ϑunif is uniform on Θ, we conclude that !d u unif , 0 < u ≤ Mopt . (C.27) Pr Cw (ϑ ) − min Cw (ϑ) ≤ u ≥ ϑ∈Θ Mopt Apply this bound with w = wλ⋆ and u = δloc /2 to obtain
Pr ∆λ,m (ϑ
unif
δ ) ≤ loc 2
≥
δloc 2Mopt
!d
.
On this event and when ∆λ,m (θ ) > δloc , we have F (θ ) − F (ϑunif ) ≥ ∆λ,m (θ ) − δloc /2 ≥ ∆λ,m (θ )/2. Since F (y (θ )) ≤ F (θ ), the two-candidate gain is nonnegative and at least F (θ ) − F (ϑunif ) for every draw. It therefore dominates the positive part below: h
i
h
Eϑunif F (θ ) − min{F (y (θ )), F (ϑunif )} ≥ Eϑunif [F (θ ) − F (ϑunif )]+ δloc 2Mopt
1 ≥ 2
!d
i
1 ∆λ,m (θ ) ≥ 2Mopt
δloc 2Mopt
!d
[∆λ,m (θ )]2 .
The last inequality uses (C.23). For the small sublevel set, the minimum of the two costs is at most F (y (θ )), so the projected-step bounds already proved apply directly to the same left-hand side. Taking 1 1 c2loc κopt := min , , 4Mopt 8αstu (Lst + 1/αstu )2 2Mopt
δloc 2Mopt
!d
>0
therefore proves (B.19) for every θ ∈ Θ. All constants were chosen using the fixed objective and parameter set, and hence are independent of t. Proof of Lemma B.10. We study the fixed true cost F (θ ) := Cwλ⋆ (θ ), so that ∆λ,m (θ ) = F (θ ) −
† F (θλ,m ). The algorithm evaluates Ct , whereas F is used only in this proof. Recall from (C.23) that 0 ≤ ∆λ,m (θ ) ≤ Mopt . For target update t ≥ 1, let Gt := Ft−1 ∨ σ (wt ). The current student θt−1 is fixed conditional on Gt . Compare the actual gradient candidate with the population candidate
ϑt,1 = ProjΘ θt−1 − αstu gbtstu , y (θt−1 ) = ProjΘ (θt−1 − αstu ∇θ F (θt−1 )) . The map y is the proof-only update in Lemma B.9. We first write the one-step error comparison, before bounding its perturbation terms. Lemma B.8 and (B.12) imply, on every sample path, F (θt ) ≤ Ct (θt ) + sup |F (θ ) − Ct (θ )| ≤ min Ct (ϑt,k ) + ξt + sup |F (θ ) − Ct (θ )| k∈{1,2}
θ∈Θ
θ∈Θ
≤ min F (ϑt,k ) + ξt + 2 sup |F (θ ) − Ct (θ )| ≤ min{F (ϑt,1 ), F (ϑt,2 )} + ξt + 4λH∥wt − wλ⋆ ∥2 . k∈{1,2}
θ∈Θ
56
Replacing one entry of a minimum changes that minimum by at most the absolute change in the entry. Since F is 4λBH 3/2 -Lipschitz according to Lemma B.5, we thus have min{F (ϑt,1 ), F (ϑt,2 )} ≤ min{F (y (θt−1 )), F (ϑt,2 )} + 4λBH 3/2 ∥ϑt,1 − y (θt−1 )∥2 . Conditional on Gt , the exploration candidate ϑt,2 is uniform on Θ. By Lemma B.9, we have E[F (θt−1 ) − min{F (y (θt−1 )), F (ϑt,2 )} | Gt ] ≥ κopt [∆λ,m (θt−1 )]2 . Since F (θt−1 ) is measurable with respect to Gt , rearrange and we have E[min{F (y (θt−1 )), F (ϑt,2 )}| Gt ] ≤ F (θt−1 ) − κopt [∆λ,m (θt−1 )]2 Thus, utilizing the inequalities we have proved above, we have E[∆λ,m (θt ) | Gt ] † =E[F (θt ) | Gt ] − F (θλ,m ) † ≤E[min{F (y (θt−1 )), F (ϑt,2 )} | Gt ] − F (θλ,m ) + 4λBH 3/2 E[∥ϑt,1 − y (θt−1 )∥2 | Gt ] + 4λH∥wt − wλ⋆ ∥2 + E[ξt | Gt ] † ≤F (θt−1 ) − F (θλ,m ) − κopt [∆λ,m (θt−1 )]2 + 4λBH 3/2 E[∥ϑt,1 − y (θt−1 )∥2 | Gt ] + 4λH∥wt − wλ⋆ ∥2 + E[ξt | Gt ]. (C.28)
We next bound the candidate discrepancy in this recursion. Nonexpansiveness of Euclidean projection gives h
i
∥ϑt,1 − y (θt−1 )∥2 ≤ αstu ∥gbtstu − ∇θ F (θt−1 )∥2 ≤ αstu ∥gbtstu − ∇θ Ct (θt−1 )∥2 + ∥∇θ Ct (θt−1 ) − ∇θ F (θt−1 )∥2 . Conditional Cauchy-Schwarz and Lemma B.6 bound the first term by h
i
h
E ∥gbtstu − ∇θ Ct (θt−1 )∥2 Gt ≤ E ∥gbtstu − ∇θ Ct (θt−1 )∥22 Gt
i1/2
≤
4λBH 3/2 √ . bt
Lemma B.7 bounds the second term by 2λH 3/2 ∥wt − wλ⋆ ∥2 . Multiplying by the cost Lipschitz constant therefore gives 4λBH 3/2 E[∥ϑt,1 − y (θt−1 )∥2 | Gt ] ≤
16αstu λ2 B 2 H 3 √ + 8αstu λ2 BH 3 ∥wt − wλ⋆ ∥2 . bt
For validation, Gt ⊆ Ht in Lemma B.8. The tower property yields E[ξt | Gt ] = E[E[ξt | Ht ] | Gt ] ≤
16λBH . t+1
To combine these bounds, write et := E[∆λ,m (θt )]. Taking total expectations in (C.28), using bt = t + 1, and substituting the two perturbation bounds gives h
i
et ≤ et−1 − κopt E [∆λ,m (θt−1 )]2 +
16αstu λ2 B 2 H 3 16λBH √ + (8αstu λ2 BH 3 + 4λH )E[∥wt − wλ⋆ ∥2 ] + . t+1 t+1
57
Applying the calibration bound (B.7) and the Cauchy-Schwarz inequality, we have
E[∥wt − wλ⋆ ∥2 ] ≤ E[∥wt − wλ⋆ ∥22 ]
1/2
2 2 ≤ √ ≤ √ . γ t+2 γ t+1
Also, E[[∆λ,m (θt−1 )]2 ] ≥ e2t−1 by Jensen’s inequality, and (t + 1)−1 ≤ (t + 1)−1/2 . Consequently, Eopt et ≤ et−1 − κopt e2t−1 + √ , t ≥ 1. t+1
(C.29)
We solve (C.29) by induction with the deterministic bound vt := Kopt (t + 1)−1/4 . Initially, we have e0 ≤ Mopt ≤ Kopt = v0 . Suppose et ≤ vt . If vt+1 ≥ Mopt , the bound et+1 ≤ Mopt proves the next step. Otherwise, we have
vt = vt+1
t+2 t+1
1/4
< 21/4 Mopt < 2Mopt .
Since κopt ≤ 1/(4Mopt ), the derivative of u − κopt u2 on [0, 2Mopt ] is 1 − 2κopt u ≥ 0. Apply this monotonicity to the recursion at time t + 1 and we have 2 −E κopt Kopt Eopt opt √ et+1 ≤ vt − κopt vt2 + √ ≤ vt − . t+2 t+1 κopt K 2
K
opt 2 −E By the definition of Kopt , we have κopt Kopt ≥ 4opt . The first inequality uses opt ≥ 2 2 Kopt ≥ 2Eopt /κopt ; the second uses Kopt ≥ 1/(2κopt ). On the other hand, integration of the derivative of u−1/4 gives
vt − v t + 1 =
Kopt 4
Z t+2
u−5/4 du ≤
t+1
Kopt Kopt . ≤ √ 4(t + 1)5/4 4 t+1
Combining the last three displays yields et+1 ≤ vt+1 . By induction, we prove the first bound in (B.20) for every T ≥ 1. Finally, we translate the true-cost bound to the calibrated optimization error. By the elementary inequality | min f − min g| ≤ sup |f − g|, applied on Θ, we have |εT − ∆λ,m (θT )| = CT (θT ) − F (θT ) + min F (θ ) − min CT (θ ) θ∈Θ
θ∈Θ
≤ 2 sup |CT (θ ) − F (θ )| ≤ 4λH∥wT − wλ⋆ ∥2 .
(C.30)
θ∈Θ
Taking expectations and applying (B.7) proves the second bound in (B.20). We finish the proof. Proof of Lemma B.11. We will apply the nonnegative form of Lemma A.1 with K = Θ, f (θ ) = ∆λ,m (θ ), g (θ ) = Kλ,m (θ ). Accordingly, we verify that the two functions are nonnegative, that their graphs are subanalytic and compact, and that ∆λ,m (θ ) = 0 implies Kλ,m (θ ) = 0. We also establish an upper bound on Kλ,m to use in the theorem’s exponent normalization. Their definitions are recalled in the lemma statement. 58
† Nonnegativity. By the optimality of θλ,m and the nonnegativity of KL divergence, respectively, we have ∆λ,m (θ ) ≥ 0, Kλ,m (θ ) ≥ 0, θ ∈ Θ.
Analyticity. Lemma B.7 proves analyticity for the average KL against any fixed positive answer ej ) and with Qj = πstu,θ† (· | x ej ), respectively. Thus laws. Apply that result with Qj = πwλ⋆ (· | x λ,m
† Cwλ⋆ and Kλ,m are analytic. Subtracting the constant Cwλ⋆ (θλ,m ) shows that ∆λ,m is analytic.
Compact subanalytic graphs. For F ∈ {∆λ,m , Kλ,m }, the restricted graph is graph(F |Θ ) = {(θ, u) ∈ Rd+1 : B 2 − ∥θ∥22 ≥ 0, u − F (θ ) = 0}. The defining inequality is polynomial and the equality is analytic by Lemma B.7. Hence the graph is semianalytic, and therefore subanalytic. It is also compact because it is the image of the compact ball Θ under the continuous map θ 7→ (θ, F (θ )). Zero-set inclusion. If ∆λ,m (θ ) = 0, its definition gives † Cwλ⋆ (θ ) = Cwλ⋆ (θλ,m ) = min Cwλ⋆ (ϑ). ϑ∈Θ
Thus θ is also a global minimizer. By Assumption 3.3, we have ej ) = πstu,θ† (· | x ej ), j = 1, . . . , m. πstu,θ (· | x λ,m
Each KL summand defining Kλ,m (θ ) is consequently the divergence of a law from itself, hence zero. We obtain the required inclusion {θ ∈ Θ : ∆λ,m (θ ) = 0} ⊆ {θ ∈ Θ : Kλ,m (θ ) = 0}. A uniform bound for the KL divergence Apply the student probability envelope (C.14) from † Lemma B.4 to θ and θλ,m . Dividing the lower bound for the numerator by the upper bound for the denominator, and conversely, gives e−4B ≤
πstu,θ (a | s) ≤ e4B (a ∈ B (s)). πstu,θ† (a | s) λ,m
ej ), put sj,h = Taking logarithms gives an absolute token log-ratio bound of 4B. For a1:H ∈ A(x (xej , a1:h−1 ). By (3.1) and the triangle inequality, we have
log
H H X X ej ) πstu,θ (a1:H | x πstu,θ (ah | sj,h ) πstu,θ (ah | sj,h ) = ≤ log ≤ 4BH. log ej ) πstu,θ† (a1:H | x πstu,θ† (ah | sj,h ) πstu,θ† (ah | sj,h ) h=1 h=1 λ,m
λ,m
λ,m
Using the expectation form of KL in the definition of Kλ,m , this pointwise bound gives 0 ≤ Kλ,m (θ ) ≤
m 1 X E [4BH ] = 4BH, m j =1 a1:H ∼πstu,θ (·|exj )
Via the steps above, we verify the nonnegativity, compact subanalytic graphs, and zero-set inclusion required by Lemma A.1. Moreover, the last step provides the bound Kλ,m (θ ) ≤ 4BH ≤ Mλ,m with 59
Mλ,m := max{1, 4BH}. Since the functions and their domain are fixed independently of the iteration number, applying Lemma A.1 and we get that there exist aλ,m > 0 and pλ,m ≥ 1 such that ∆λ,m (θ ) ≥ aλ,m [Kλ,m (θ )]pλ,m θ ∈ Θ, We finish the proof.
D
Proofs in Section 6
Proof of Proposition 6.1. We first compute the true student objective. For either prompt in pair (xei,+ , xei,− ), we have ei,ϵ , ∅) = πstu,θ (a | x
exp(θi 1{a = 1}) , a ∈ {0, 1, null}. eθi + 2
θ
We define p(θi ) := eθei +i 2 = σ (θi − log 2), σ (u) = (1 + e−u )−1 . Then, the probabilities of the student policy to output 0 and null are each [1 − p(θi )]/2. Therefore, we have 1 X 1 ei,ϵ , a1:2 )] = [πstu,θ (1 | x ei,+ , ∅) + πstu,θ (0 | x ei,− , ∅)] Ea1:2 ∼πstu,θ (·|exi,ϵ ) [R(x 2 ϵ∈{+,−} 2 1 − p ( θi ) 1 + p ( θi ) 1 = . = p ( θi ) + 2 2 4
Both the student and reference emit EOS with probability one at h = 2, therefore, for ϵ ∈ {+, −}, ei,ϵ , we can explicitly compute the KL divergence as at any x p ( θi ) 1 − p ( θi ) [1 − p(θi )]/2 +2 log 1/3 2 1/3 1 − p ( θi ) p ( θi ) + [1 − p(θi )] log = p(θi ) log 1/3 2/3 θ i e +2 = θi p(θi ) − log . 3
ei,ϵ )∥πpre (·|x ei,ϵ )) = p(θi ) log KL(πstu,θ (·|x
Recall that p′ (θi ) = p(θi )[1 − p(θi )]. Hence, for each i ∈ [d] and ϵ ∈ {+, −}, we have " ei,ϵ ) ∥ πpre (· | x ei,ϵ )) =∇θ ∇θ KL(πstu,θ (· | x
eθi + 2 θi p(θi ) − log 3
"
#
#
eθi = p(θi ) + θi p (θi ) − θi ei e +2 ′
= θ i p′ ( θ i ) e i . Here ei appears because the expression depends on θ only through its ith coordinate.
60
(D.1)
ei,+ and x ei,− , we obtain Thus, using the fact that the KL regularization terms are equal at x d 1X 1 + p ( θi ) ei,+ ) ∥ πpre (· | x ei,+ )) . Jλ,m (πstu,θ ) = − λKL(πstu,θ (· | x d i=1 4
Differentiating this finite sum and substituting the preceding KL-gradient identity gives d p′ (θi ) 1X ei,+ ) ∥ πpre (· | x ei,+ )) ei − λ∇θ KL(πstu,θ (· | x d i=1 4
∇θ Jλ,m (πstu,θ ) =
d 1X 1 p′ ( θ i ) − λθi ei d i=1 4
=
p1 (θ1 )[1 − p1 (θ1 )] 41 − λθ1 1 .. = . . d pd (θd )[1 − pd (θd )] 14 − λθd Each summand strictly increases up to 1/√ (4λ) and strictly decreases afterward. The vector of these maximizers is feasible and interior, since d/(4λ) < B. Then, we apply the first order optimality † 1 condition to prove that θλ,m = 4λ 1d ∈ int(Θ). We next identify the unrestricted optimal policy and verify that it belongs to the teacher policy class but not to the student policy class. Specifically, by Lemma B.1, we have that πλ⋆ ((a1 , EOS) | x) =
e1{a1 =y (x)}/λ (1/3)e1{a1 =y (x)}/λ = = π(√2/λ)u1 ((a1 , EOS) | x). (e1/λ + 2)/3 e1/λ + 2
√ For wλ⋆ = ( 2/λ)u1 , by the feature definition in teacher class, we have (
(wλ⋆ )⊤ ϕ((x, ∅), a1 ) =
1/λ, a1 = y (x), 0, a1 ̸ = y ( x ) .
The resulting softmax probabilities are exactly those in the preceding expression. Since ∥wλ⋆ ∥2 = √ 2/λ < B, this proves that wλ⋆ ∈ W and πλ⋆ = πwλ⋆ . To prove that no student represents this policy, we compare the probability of the answer (1, EOS) ei,+ and x ei,− . The unrestricted optimum satisfies at the two prompts x ei,+ ) = πλ⋆ ((1, EOS) | x
e1/λ 1 > 1/λ = πλ⋆ ((1, EOS) | xei,− ). 1/λ e +2 e +2
In contrast, every student satisfies ei,+ ) = p(θi ) = πstu,θ ((1, EOS) | x ei,− ). πstu,θ ((1, EOS) | x
Thus no θ ∈ Θ can match the unrestricted optimum at both prompts. It remains to compare the teacher’s regularized return with that of every student. We first obtain a student upper bound using the oracle parameter already proved optimal.
61
† Substituting (θλ,m )i = 1/(4λ) into the objective and the KL expression in (D.1) gives d 1X 1 + p(1/(4λ)) p(1/(4λ)) e1/(4λ) + 2 Jλ,m (πstu,θ ) ≤ Jλ,m (πstu,θ† ) = −λ − log λ,m d i=1 4 4λ 3
"
!#
d 1X 1 e1/(4λ) + 2 1 e1/(4λ) + 2 + λ log = + λ log . = d i=1 4 3 4 3
"
#
We now compute the teacher’s regularized return. Since wtea = αwλ⋆ , its answer probabilities are πtea ((a1 , EOS) | x) =
exp(α1{a1 = y (x)}/λ) . eα/λ + 2
In particular, its expected reward at every target prompt is Ea1:2 ∼πtea (·|x) [R(x, a1:2 )] = πtea ((y (x), EOS) | x) =
eα/λ . eα/λ + 2
Because α < 1, this correct-verdict probability is strictly smaller than that of πλ⋆ . Hence the teacher differs from the unrestricted optimal policy. The reference assigns probability 1/3 to each feasible answer. Consequently, we know that πtea (a1:2 | x) exp(αR(x, a1:2 )/λ) α eα/λ + 2 log = log · 3 = R ( x, a ) − log . 1:2 πpre (a1:2 | x) λ 3 eα/λ + 2
(D.2)
Substituting this identity into the regularized objective yields m eα/λ + 2 1 X ej , a1:2 ) + λ log Ea1:2 ∼πtea (·|exj ) (1 − α)R(x Jλ,m (πtea ) = m j =1 3
"
#
eα/λ eα/λ + 2 = (1 − α) α/λ + λ log . 3 e +2 Viewing this expression as a function of α, differentiation implies that d eα/λ 2(1 − α)eα/λ eα/λ 2(1 − α)eα/λ Jλ,m (πtea ) = − α/λ + + = > 0. dα e + 2 λ(eα/λ + 2)2 eα/λ + 2 λ(eα/λ + 2)2 1/(2λ)
1/(2λ) +2
e e Thus, for every α ∈ [1/2, 1), we have that Jλ,m (πtea ) ≥ 12 e1/ (2λ) +2 + λ log
3
.
We finally compare this teacher lower bound with the student upper bound. Since es > 1 for s > 0, we have es /(es + 2) > 1/3. Hence by direct algebra, we have e1/(2λ) 1 e1/(2λ) + 2 > , log = 3 e1/(2λ) + 2 e1/(4λ) + 2
Z 1/(2λ) 1/(4λ)
es 1 ds > s e +2 3
1 1 − 2λ 4λ
=
1 . 12λ
Combining these two inequalities with the preceding bounds, we obtain Jλ,m (πtea ) − Jλ,m (πstu,θ† ) ≥ λ,m
1 e1/(2λ) 1 e1/(2λ) + 2 1 1 1 − + λ log > − +λ = 0. 1/ ( 2λ ) 1/ ( 4λ ) 2e 12λ +2 4 e +2 6 4
62
Therefore, we have Jλ,m (πtea ) > Jλ,m (πstu,θ† ) ≥ Jλ,m (πstu,θ ) for every θ ∈ Θ. We finish the λ,m
proof. Proof of Theorem 6.2. We first identify the minimizer of the direct-matching objective and its distance from the oracle student. Then, we bound the SGD error around this minimizer. Finally, a quadratic lower bound for the policy KL divergence converts the remaining parameter distance into the claimed separation. Recall from the proof of Proposition 6.1 that p(θi ) = eθi /(eθi + 2). For a target prompt x and a feasible answer a1:2 , all three policy probabilities are positive. Factoring the likelihood ratio through the reference gives "
#
πstu,θ (a1:2 | x) πpre (a1:2 | x) π (a1:2 | x) ZSM (θ; x, a1:2 ) = log + λ log stu,θ πpre (a1:2 | x) πtea (a1:2 | x) πpre (a1:2 | x)
= (1 + λ) log
πstu,θ (a1:2 | x) πtea (a1:2 | x) − log πpre (a1:2 | x) πpre (a1:2 | x)
= (1 + λ) log
πstu,θ (a1:2 | x) α eα/λ + 2 − R(x, a1:2 ) + log . πpre (a1:2 | x) λ 3
The second equality uses log(uv ) = log u + log v and log(1/u) = − log u; the last uses (D.2). We now derive the population cost from its definition. Using (6.11), the definition of ZSM , and its pointwise expansion above, we obtain CSM (θ ) =
=
m 1 X
m j =1
KL(πstu,θ (· | x ej ) ∥ πtea (· | x ej )) + λKL(πstu,θ (· | x ej ) ∥ πpre (· | x ej ))
m 1 X E [ZSM (θ; xej , a1:2 )] m j =1 a1:2 ∼πstu,θ (·|exj )
m ej ) πstu,θ (a1:2 | x 1 X α eα/λ + 2 ej , a1:2 ) + log = Ea1:2 ∼πstu,θ (·|exj ) (1 + λ) log − R (x ej ) m j =1 πpre (a1:2 | x λ 3
"
=
m 1 X
m j =1
(1 + λ)KL(πstu,θ (· | x ej ) ∥ πpre (· | x ej )) −
#
α
α eλ + 2 ej , a1:2 )] + log Ea1:2 ∼πstu,θ (·|exj ) [R(x . λ 3
The last equality uses linearity of expectation and the definition of KL. The final logarithm is independent of both the answer and the prompt, so averaging leaves it unchanged. Since m = 2d, we can replace the sum over j by the sum over i ∈ [d] and ϵ ∈ {+, −}. The calculations in Proposition 6.1 give X
ei,ϵ ) ∥ πpre (· | x ei,ϵ )) = 2KL(πstu,θ (· | x ei,+ ) ∥ πpre (· | x ei,+ )) , KL(πstu,θ (· | x
ϵ∈{+,−}
X ϵ∈{+,−}
ei,ϵ , a1:2 )] = 2 Ea1:2 ∼πstu,θ (·|exi,ϵ ) [R(x
63
1 + p ( θi ) 1 + p ( θi ) = . 4 2
Substituting these two identities into the preceding cost expression gives d eα/λ + 2 1 X α 1 + p ( θi ) ei,+ ) ∥ πpre (· | x ei,+ )) − + log CSM (θ ) = 2(1 + λ)KL(πstu,θ (· | x 2d i=1 λ 2 3
d 1X α eα/λ + 2 = (1 + λ)KL(πstu,θ (· | xei,+ ) ∥ πpre (· | xei,+ )) − [1 + p(θi )] + log . d i=1 4λ 3
We next differentiate this expression with respect to the full parameter vector θ. Notice that the final logarithm is constant in θ, and ∇θ p(θi ) = p′ (θi )ei . Using the KL-gradient identity already proved in Proposition 6.1, we obtain d 1X α (1 + λ)∇θ KL(πstu,θ (· | xei,+ ) ∥ πpre (· | xei,+ )) − ∇θ [1 + p(θi )] d i=1 4λ
∇θ CSM (θ ) =
d d 1X α α 1X ei . p′ (θi ) (1 + λ)θi − (1 + λ)θi p′ (θi )ei − p′ (θi )ei = d i=1 4λ d i=1 4λ
=
α ⋆ : = argmin By the first order optimality condition, the minimizer is θSM θ∈Θ CSM (θ ) = 4λ(1+λ) 1d . √ ⋆ ∥ < It is feasible and interior because ∥θSM d/(4λ) < B. Using the oracle parameter from 2 Proposition 6.1, we obtain the fixed parameter gap √ √ d d(1 − α ) α † ⋆ ∥θSM − θλ,m ∥2 = . (D.3) 1− ≥ 4λ 1+λ 4λ
The inequality follows from α/(1 + λ) ≤ α. ⋆ , starting from the update itself. Define We now bound the error of the SGD iterate relative to θSM SM is F SM -measurable, and θ ⋆ is a fixed point FtSM := σ jsSM , aSM t s,1:2 : 0 ≤ s < t . The parameter θt SM of the projection onto Θ. Therefore, the update and nonexpansiveness of projection yields that ⋆ 2 SM SM SM ⋆ bt ) − ProjΘ (θSM ∥θtSM ) +1 − θSM ∥2 = ProjΘ (θt − ηt g
2
2 SM ⋆ SM SM 2 ≤ ∥θt − θSM − ηt gbt ∥2 ⋆ ⋆ = ∥θtSM − θSM ∥22 − 2ηtSM (θtSM − θSM )⊤ gbtSM + (ηtSM )2 ∥gbtSM ∥22 .
Taking conditional expectations on both sides, we have that h
i
h
i
h
i
⋆ 2 SM ⋆ ⋆ E ∥θtSM ≤ ∥θtSM − θSM ∥22 − 2ηtSM E (θtSM − θSM )⊤ gbtSM FtSM + (ηtSM )2 E ∥gbtSM ∥22 FtSM . +1 − θSM ∥2 Ft (D.4) Since ηtSM > 0, an upper bound for the next error requires a lower bound for the cross term and an upper bound for the gradient second moment. We establish these two bounds in turn.
ei,ϵ , the student score is For the cross term, we first identify the conditional mean of gbtSM . At x = x h
i
Sstu,θ (x, a1:2 ) = ∇θ θi 1{a1 = 1} − log(eθi + 2) = ei [1{a1 = 1} − p(θi )]. Thus, we obtain ∥Sstu,θ (x, a1:2 )∥2 ≤ 1, Ea1:2 ∼πstu,θ (·|x) [Sstu,θ (x, a1:2 )] = ei [p(θi ) − p(θi )] = 0.
64
From the definitions of the score and sampled cost, we have ∇θ πstu,θ (a1:2 | x) = πstu,θ (a1:2 | x)Sstu,θ (x, a1:2 ), ∇θ ZSM (θ; x, a1:2 ) = (1 + λ)Sstu,θ (x, a1:2 ). The product rule for the finite sum defining CSM , followed by the zero-mean score identity, gives ∇θ CSM (θ ) =
=
m h i 1 X ej , a1:2 ) ZSM (θ; x ej , a1:2 ) + 1 + λ Ea1:2 ∼πstu,θ (·|exj ) Sstu,θ (x m j =1 m 1 X E [Sstu,θ (xej , a1:2 )ZSM (θ; xej , a1:2 )] . m j =1 a1:2 ∼πstu,θ (·|exj )
Since the algorithm samples jtSM uniformly and then samples the answer from the current student,
SM Pr jtSM = j, aSM = t,1:2 = a1:2 Ft
1 ej ), a1:2 ∈ A(x ej ) . π SM (a1:2 | x m stu,θt
Thus the preceding gradient identity implies E[gbtSM | FtSM ] = ∇θ CSM (θ )|θ =θSM .
(D.5)
t
We now use this unbiasedness identity to bound the cross term. For every u ∈ [−B, B ], by algebra, we have p′ (u) = σ ′ (u − log 2) ≥ σ ′ (B + log 2) = κSM , ⋆ ) = α/ (4λ), we Using the population gradient already computed above, together with (1 + λ)(θSM i have i h ⋆ ⋆ )⊤ E[gbtSM | FtSM ] E (θtSM − θSM )⊤ gbtSM FtSM =(θtSM − θSM ⋆ =(θtSM − θSM )⊤ ∇θ CSM (θ )|θ=θSM t
=
d X
1+λ p′ (θtSM )i d i=1
h
⋆ (θtSM )i − (θSM )i
(D.6)
i2
(1 + λ)κSM SM ⋆ ⋆ ∥θt − θSM ∥22 = µSM ∥θtSM − θSM ∥22 . d The first equality uses the measurability of θtSM ; the second uses (D.5). ≥
For the second moment in (D.4), the definition of the gradient estimate and the score bound already proved give h
E ∥gbtSM ∥22
FtSM
i
=E
ej SM , aSM Sstu,θSM (x t,1:2 ) t t
ej SM , aSM ZSM (θtSM ; x t,1:2 ) t
≤E
2 2 2
ej SM , aSM ZSM (θtSM ; x t,1:2 ) t
FtSM
2
FtSM
.
It therefore suffices to bound the sampled cost uniformly. The explicit student/reference ratio is log
ei,ϵ ) πstu,θ (a1:2 | x eθi + 2 = θi 1{a1 = 1} − log . ei,ϵ ) πpre (a1:2 | x 3
65
The average (eθi + 1 + 1)/3 lies between emin{0,θi } and emax{0,θi } . Hence both terms on the right lie between min{0, θi } and max{0, θi }, and the absolute log ratio is at most |θi |. Similarly, (D.2) tea (a1:2 |x) α and 0 ≤ log[(eα/λ + 2)/3] ≤ α/λ yields log ππpre (a1:2 |x) ≤ λ . Returning to the reference decomposition of ZSM , we obtain πstu,θ (a1:2 | x) πtea (a1:2 | x) + log πpre (a1:2 | x) πpre (a1:2 | x) α ≤ (1 + λ)|θi | + ≤ (2 + λ)B = GSM . λ
|ZSM (θ; x, a1:2 )| ≤ (1 + λ) log
ei,ϵ and α/λ ≤ B. This uniform bound proves the required second-moment inequality: Here x = x h
i
E ∥gbtSM ∥22 FtSM ≤ G2SM .
(D.7)
Substituting (D.6) and (D.7) into the original error recurrence (D.4), we obtain h
i
⋆ 2 SM ⋆ E ∥θtSM ≤(1 − 2µSM ηtSM )∥θtSM − θSM ∥22 + (ηtSM )2 G2SM +1 − θSM ∥2 Ft
=
t G2 ⋆ ∥θtSM − θSM ∥22 + 2 SM 2 , t+2 µSM (t + 2)
Taking expectations and applying the law of total expectation, we have that h
i
⋆ 2 E ∥θtSM +1 − θSM ∥2 ≤
h i t G2 ⋆ E ∥θtSM − θSM ∥22 + 2 SM 2 . t+2 µSM (t + 2)
(D.8)
We solve this recurrence by induction to obtain h
i
⋆ E ∥θTSM − θSM ∥22 ≤
G2SM . µ2SM (T + 1)
(D.9)
⋆ ∥ ≤ B, and Indeed, the initial bound follows from θ0SM = 0, ∥θSM 2
GSM d(2 + λ)B = ≥ 4dB ≥ B, µSM (1 + λ)κSM where κSM ≤ 1/4. For the induction step, substitution of the bound at time t into the expected recurrence gives h
⋆ 2 E ∥θtSM +1 − θSM ∥2
i
t 1 G2SM 1 1 G2SM G2 = ≤ ≤ 2SM + − . µSM (t + 2)(t + 1) (t + 2)2 µ2SM t + 2 (t + 1)(t + 2)2 µ2SM (t + 2)
It remains to translate the fixed parameter gap and the SGD error into the policy KL in the
66
theorem. For any θ ∈ Θ, the explicit student policy πstu,θ gives
ei,ϵ ) ei,ϵ ) πstu,θ† (· | x KL πstu,θ (· | x λ,m
=
†
ei,ϵ ) θi − (θ † πstu,θ ((a1 , EOS) | x
X
e(θλ,m )i + 2 1{a = 1} + log ) 1 λ,m i eθi + 2
a1 ∈{0,1,null}
†
=
† θi − (θλ,m )i
† = (θλ,m ) i − θi
e(θλ,m )i + 2 p(θi ) + log e θi + 2
2 Z 1 0
† )i − θi ] ds ≥ (1 − s) p′ θi + s[(θλ,m
2 κSM † θi − (θλ,m )i . 2
The last equality is by Taylor’s integral formula for u 7→ log(eu + 2), whose first and second derivatives are p(u) and p′ (u). The segment between the two coordinates lies in [−B, B ], so the already proved lower bound p′ (u) ≥ κSM applies throughout the integral. Averaging over the 2d prompts gives m 1 X κSM † ej ) πstu,θ† (· | x ej ) ≥ KL πstu,θ (· | x ∥θ − θλ,m ∥22 . (D.10) λ,m m j =1 2d By the triangle inequality, Cauchy-Schwarz, and (D.9), we have h
i
h
i
h
† Combining this with (D.3), we obtain E ∥θTSM − θλ,m ∥22
i1/2
† † † ⋆ ⋆ ∥θSM − θλ,m ∥2 ≤ E ∥θTSM − θλ,m ∥2 + E ∥θTSM − θSM ∥2 ≤ E ∥θTSM − θλ,m ∥22
h
√
≥
i1/2
+
d(1−α) − µ G√SMT +1 4λ SM
GSM √ . µSM T + 1
. +
We can therefore start from the required policy error and conclude
m i 1 X κSM h SM † ej ) πstu,θ† (· | x ej ) ≥ E KL πstu,θSM (· | x E ∥θT − θλ,m ∥22 T λ,m m j =1 2d "√ #2 κSM GSM d(1 − α ) √ . ≥ − 2d 4λ µSM T + 1 +
The first inequality is (D.10); the second is the square of the preceding bound. For the stated iteration threshold, substituting the definition of µSM gives 64λ2 dG2SM T +1 ≥ = (1 + λ)2 κ2SM (1 − α)2
8λGSM √ µSM d(1 − α)
!2
.
√ √ Hence GSM /[µSM T + 1] ≤ d(1 − α)/(8λ), and the KL lower bound is at least κSM 2d
"√
d(1 − α ) − 4λ
√
d(1 − α ) 8λ
67
#2
=
κSM (1 − α)2 > 0. 128λ2