Skip a Layer or Loop It? Learning Program-of-Layers in LLMs
标题:跳过一层还是循环复用?大语言模型中层序列程序的学习方法
arxiv:2606.06574v1|github: https://github.com/tianyi-lab/PoLar
一、介绍
-
通用基础模型对所有的输入都同一执行一个静态的、预训练好的架构,即使输入在复杂性、难度上大相径庭。相比之下,传统的编程解决问题会在算法架构和复杂度上更加的灵活和针对性。比如,对于一个经验丰富的程序员,在处理简单任务时能够精简步骤、高效运算,同时也懂得合理调整时空复杂度,用以解决难度更高的问题。但这类程序是针对每一类问题专门设计并优化的,通用性不如大语言模型(LLMs)。这就引入一个问题:对不同任务采用相同架构或“程序”(即按照固定顺序完成所有层的前向传播)是否总能实现最优且高效的效果?通用模型能否针对每一份输入进一步优化其所使用的“程序”?
-
文中,作者将预训练的LLM每一层形式化为一个原子函数,程序可以按任意时间和顺序来调用,通过该方法,便可将推理阶段的动态模型架构表示为对于不同输入的层程序(POLAR)。作者借助蒙特卡洛树搜索(MCTS)对层程序的空间进行搜寻,研究发现:对于每一项所评估的输入任务,几乎总能找到效果更优(精度更高且/或长度更短)的程序。

-
相较之前的相关研究,本文分析MCTS得到的有效执行程序,表明在跳跃与循环的联合下,通常能够找到远优于单独一种方法所得到的程序。虽然大多数效果不错的程序能够短于默认长度,但通过跳转/循环操作提升程序复杂度,能够大幅改善输出质量,在难度更高的任务上效果尤为明显。另外,大多数成功的程序主要由连续的层段组成。
-
这些实验观察不仅证实更优的POLAR方案无需额外训练即可广泛适用,还启发提出一种实用的POLAR学习预测方法,规避了MCTS带来的高昂开销。这样的设计带来了几个好处:
- 通过消除逐个输入搜索的昂贵成本,使得层级程序推理计算可行。
- 通过在统一的执行框架下,组合层跳跃和重复,POLAR对以往单一方式的动态深度方法进行严格泛化。
- 能够在完全冻结的模型中实现灵活的测试时计算扩展,让推理过程可适配输入难度,同时维持模型的泛化能力。
-
最后用多个预训练的LLM在数学推理问题上进行评估测试。结果证实POLAR的准确率,相对于标准的推理和其他动态深度方法均得到了提高,且平均执行的层数更少。此外,增加候选执行程序能够实现显著的测试时计算规模缩放,程序能够泛化到跨不同领域的分布外的基准。
二、层级程序(POLAR)在LLM中的动态推理。
1. 核心问题
标准的前向传播是否总是最优?能否通过改变层的执行顺序(跳过/循环)来挖掘模型潜在的推理能力?
2. 将推理视作程序执行
在这种观点中,推理是一个逐步选择和组合预训练模块的过程,根据每个输入在长度和顺序上都不同,而每个模块视为一个固定的预训练函数。例如有D个transfomer层的预训练LLM,每一层定义为一个固定的计算函数:fi:RT×d→RT×d,i∈{0,…,D−1}f_i:\mathbb{R}^{T \times d} \to \mathbb{R}^{T \times d},i \in \{0,\dots,D-1\}fi:RT×d→RT×d,i∈{0,…,D−1},一个程序定义为一个有限的层索引序列:π=(i1,i2,⋯ ,iK),ik∈{0,…,D−1}\pi = (i_1,i_2,\cdots,i_K),i_k \in \{0,\dots,D-1\}π=(i1,i2,⋯,iK),ik∈{0,…,D−1},从而形成组合计算:Fπ=fiK∘⋯∘fi1F_{\pi} = f_{i_K} \circ \dots \circ f_{i_1}Fπ=fiK∘⋯∘fi1。执行程序时,便将输入应用到该组合上,如果生成正确的预测便视作程序有效valid。
3. 搜寻有效的执行程序
作者利用MCTS作为离线诊断工具,探索了有效程序的空间,得出四项关键发现::
-
发现1(互补性).对层仅重复比仅跳过执行更好,将这两个互补的操作合并能产生最好的层级程序。(Layer recurrence-only performs better than layer skipping-only, but combining the two complementary operators produces the best program-of-layers)

-
发现2(奥卡姆剃刀).大多数有效的程序通常比标准前向传播短。(Most valid execution programs are often shorter than the standard forward pass)

-
图3表明,即便整体计算量被限制在远小于标准前向传播的水平,大量输入样本依然能够被求解。

-
图4表明,在各类模型中,我们经常能够找到有效的执行方案,这类方案所需的层级调用次数少于标准推理方式。
-
在标准推理已正确求解的输入样本(C→C)中,75.5% 的样本存在更短的有效程序。即便是最初求解错误的输入样本(W→C),也有 36.2% 的样本能够找到更短程序,以此修正模型的预测结果。上述结果表明,标准推理往往存在过度计算问题,并且通常只需大幅减少潜在计算步骤,就能实现正确推理。
-
-
发现3(复杂度扩展).系统地提升潜在执行复杂度,能够拓展有效程序的搜索空间,并增强对高难度输入的推理能力。(Increasing latent execution complexity systematically expands the space of valid programs and improves inference on harder inputs.)
对于复杂输入,执行方案的复杂度会更大。对于模型和数据集,增加执行深度和架构灵活性可系统地增大有效程序空间和提升推理准确率。

- 测试时缩放计算量可拓宽潜在推理有效执行程序的空间(Test-time scaling expands the space of valid execution programs for latent reasoning),如图5a所示,分配更多的测试时计算量——通过递归——会单调增加有效执行程序的数量。即更多的计算产生更大的可行程序空间和更高的正确性。
- 更难的输入需要更复杂的执行程序(Harder inputs require more complex execution programs),如图5b所示,对于大多数模型而言,依赖层循环或跳跃机制才能求解的输入占比会随难度提升而上升。随着难度增加,可行的执行程序受到更多约束,更加依赖复杂执行结构。表明更高的潜在执行复杂度不仅起到辅助作用,在处理难度更高的输入时往往是必须的。
- 推理精度会随执行深度系统性提升(Inference accuracy improves systematically with execution depth)

如图6所示,对于所有的模型和难度层级,有效执行程序的平均准确率随着整体执行深度而增加,这解释了计算和准确率之间的权衡:更复杂的输入往往受益于——甚至常常需要——更深层次或递归式的执行,才能实现正确的推理。
-
发现4(结构偏好).有效执行程序主要由连续的层片段组成,通常每一段最多有一个循环。(Valid execution programs are predominantly composed of contiguous layer segments and typically require at most a single recurrence per segment.)

如图7所示,预训练模型找到的有效执行程序在结构上明显倾向于简洁形式。片段segment是指作为一个整体执行的一组层,这些层不必连续,循环执行即同一segment的重复运行。因此通过统计每个片段内连续层的数量来分析片段结构。绝大多数有效程序以高度局部化的片段为主,最多仅包含一处循环;远距离跳转与深度迭代复用的情况十分少见。如图7(a)显示,54.5%的片段仅包含单层结构,超过三分之二的片段最多包含两层连续结构,而由非连续层组成的片段占比不足3.2%(比如1,4,7这种跳层分段)。与之相符,图7(b)表明大多数片段最多重复一次。因此预训练模型作为执行程序生成器:其训练目标更倾向短距离、局部复用,而非丰富的程序组合与复杂控制流。
这些研究结果表明,标准推理方式仅能从海量有效的潜在程序空间中选取单一执行方案。虽然MCTS能够挖掘该程序空间,但由于搜索空间过大,难以直接用于推理任务。这促使探索新思路:不在进行搜索,而是研究轻量化模型能否直接预测执行程序。

图2对比了基于MCTS的顺序搜索方法与论文提出的程序直接预测方法。后续研究中,将采用这种基于学习的替代方案,保留MCTS所发掘的潜在程序选择优势,同时摒弃顺序搜索流程。
MCTS形式化表述:将程序发现的过程建模为一个序列决策过程,并采用蒙特卡洛树搜索(MCTS)来探索受约束的程序空间。
- 状态和行动:搜索树中的每个节点对应一个局部或完整的执行程序π;行动是对当前程序进行修改的操作。通过应用一个跳过或重复操作,生成一个新的程序节点。
- 奖励:对于一个完成的程序π和输入x(真实答案为y):r(π,x)=IFπ(x)=yr(\pi,x)=\mathbb{I}{F_{\pi}(x)=y}r(π,x)=IFπ(x)=y,得到一个二元奖励。
- 树策略:在MCTS的每次迭代中,算法从根节点(标准前向传播)出发,沿着UCB值最大的子节点向下选择,直到到达一个叶子节点,UCB公式平衡探索和利用,以及惩罚冗长的程序
UCB(π)=R(π)v(π)⏟利用项+clnVv(π)⏟探索项−λ∣π∣D⏟长度惩罚项UCB(\pi)=\underbrace{\frac{R(\pi)}{v(\pi)}}_{利用项}+\underbrace{c\sqrt{\frac{lnV}{v(\pi)}}}_{探索项}-\underbrace{\lambda\frac{|\pi|}{D}}_{长度惩罚项}UCB(π)=利用项v(π)R(π)+探索项cv(π)lnV−长度惩罚项λD∣π∣
其中R(π)R(\pi)R(π)为累积奖励,v(π)v(\pi)v(π)为访问次数,VVV为总模拟次数,λ\lambdaλ对冗长程序施加惩罚。
三、通过LLM学习层级程序(POLAR)
基于前面得分析,本文提出POLAR方法,该方法可在推理阶段对预训练语言模型进行编程,生成适配输入的执行程序(图2)。POLAR能够动态划分预训练层并将其组合为可复用模块,在无需更新参数的前提下实现灵活计算。
3.1 程序表示
分段:将预训练模型的D层划分为若干连续分段:[0=s1,s2),[s2,s3),……,[sM,sM+1=D)[0 = s_1, s_2),[s_2, s_3),……,[s_M, s_{M+1} = D)[0=s1,s2),[s2,s3),……,[sM,sM+1=D)且每个分段长度满足约束 sj+1−sj≤Kmaxs_{j+1} - s_j \le K_{max}sj+1−sj≤Kmax(默认Kmax=4K_{max}=4Kmax=4)。分段由二进制边界掩码 zseg(x)∈{0,1}Dz^{seg}(x) \in \{0, 1\}^Dzseg(x)∈{0,1}D 表示,其中若 ziseg=1z^{seg}_i=1ziseg=1,代表第i层为一个新分段的起始位置。
操作:对于每个片段[sj,sj+1][s_j,s_{j+1}][sj,sj+1],可选择以下3个操作之一(跳过、保留、重复)来定义片段如何执行:skip:∅keep:[sj,…,sj+1−1],repeat:[sj,…,sj+1−1,sj,…,sj+1−1]skip:∅\\ keep:[s_j,\dots,s_{j+1}-1],\\ repeat:[s_j,\dots,s_{j+1}-1,s_j,\dots,s_{j+1}-1]skip:∅keep:[sj,…,sj+1−1],repeat:[sj,…,sj+1−1,sj,…,sj+1−1]操作由类别标签向量zop(x)∈{跳过,保留,重复}D\boldsymbol{z}^{op}(x) \in \{\text{跳过},\text{保留},\text{重复}\}^Dzop(x)∈{跳过,保留,重复}D 表示,其中仅当ziseg=1\boldsymbol{z}^{seg}_i=1ziseg=1(即在片段起始位置)时ziop\boldsymbol{z}^{op}_iziop才有定义;其余所有位置上的标签均会被忽略。重复操作在原理上并不局限于单次循环,可拓展为{repeat-2,……,repeat-k},以支持单个片段内的多次循环执行。
3.2 层级程序(POLAR)预测网络
根据3.1中的程序表示定义,作者对输出训练了一个轻量级的预测器。
-
架构:
- 给定输入x,首先使用嵌入模型(Qwen3-Embedding-0.6B)对其进行编码,得到token级表征 H=E(x)∈RT×dqH = E(x) \in \mathbb{R}^{T \times d_q}H=E(x)∈RT×dq其中T为token长度,dqd_qdq是该嵌入模型的隐藏层维度。随后将token表征映射至工作维度d:H~=HWh∈RT×d\tilde{H} = HW_h \in \mathbb{R}^{T \times d}H~=HWh∈RT×d
- 层级查询:为每个预训练Transformer层分配一个可学习嵌入向量ei∈Rd\boldsymbol{e}_i \in \mathbb{R}^dei∈Rd,并将其堆叠构成矩阵E∈RD×d\boldsymbol{E} \in \mathbb{R}^{D\times d}E∈RD×d。这些嵌入向量充当各层级专属查询向量。
- 交叉注意力机制:采用多头交叉注意力机制,以层嵌入作为查询向量,token嵌入作为键/值向量:X=MHA(Q,K,V),Q=E,K=H~,V=H~X = \text{MHA}(Q, K, V),Q = E,K = \tilde{H},V = \tilde{H}X=MHA(Q,K,V),Q=E,K=H~,V=H~,输出 X∈RD×d\boldsymbol{X} \in \mathbb{R}^{D\times d}X∈RD×d 为每一层生成基于输入的表征。
- 跨层编码器:层的决策不是独立的,为对模型在深度上存在的依赖关系进行建模,采用轻量级Transformer编码器沿着层级维度进行自注意力计算:X′=ENClayer(X)∈RD×dX'=ENC_{layer}(X) \in \mathbb{R}^{D \times d}X′=ENClayer(X)∈RD×d使得每一层的决策都能够依托全局深度上下文信息。
- 预测头部:对于模型的每一层,网络输出两个预测结果(分割logits、操作logits):ℓseg=X′Wseg+bseg∈RD\ell_{\text{seg}} = X'W_{\text{seg}}+b_{\text{seg}} \in \mathbb{R}^Dℓseg=X′Wseg+bseg∈RDℓop=X′Wop+bop∈RD×3\ell_{\text{op}} = X'W_{\text{op}}+b_{\text{op}} \in \mathbb{R}^{D\times3}ℓop=X′Wop+bop∈RD×3
-
监督信号来自有效执行程序:监督信号来自通过MCTS离线搜索到的“有效执行程序”。每个程序都会被解析为程序表示,按照前面定义的格式生成真实分割标签与操作标签zseg(x),zop(x)z^{seg}(x),z^{op}(x)zseg(x),zop(x)。若某输入存在多条有效程序,且其中至少一条的长度小于模型深度,则降低全深度执行对应的损失权重。该设计依据发现2,即在保留原始计算监督信息的前提下,模型更倾向于使用更短的有效程序。
-
训练目标:训练预测器,对于一批输入,收集MCTS找到的有效程序。将这些程序被解析成前面定义的“程序表示”,即分段标签与操作标签 (zseg∗(x),zop∗(x))(\mathbf{z}^{seg*}(x), \mathbf{z}^{op*}(x))(zseg∗(x),zop∗(x))。令 piseg=σ(ℓiseg)p_i^{\text{seg}} = \sigma(\ell_i^{\text{seg}})piseg=σ(ℓiseg),piop=SOFTMAX(ℓiop)\mathbf{p}_i^{\text{op}} = \text{SOFTMAX}(\ell_i^{\text{op}})piop=SOFTMAX(ℓiop)。
模型在第i层输出两个原始预测值(logits):ℓiseg\ell_i^{\text{seg}}ℓiseg(一个标量,表示是否分段)和 ℓiop\ell_i^{\text{op}}ℓiop(一个向量,包含三个数值,分别对应 skip/keep/repeat 三种操作)。
- Sigmoid 函数 σ(⋅)\sigma(\cdot)σ(⋅):用于二分类问题。将任意实数 ℓ\ellℓ 映射到 (0,1)(0, 1)(0,1) 区间。piseg=σ(ℓiseg)p_i^{\text{seg}} = \sigma(\ell_i^{\text{seg}})piseg=σ(ℓiseg) 表示第i层作为“分段起点”的预测概率。例如,piseg=0.9p_i^{\text{seg}}=0.9piseg=0.9 表示模型认为这里极有可能是新分段的开始。
- Softmax 函数:用于多分类问题。它将一个向量映射成概率分布,使得所有元素之和为1。piop=SOFTMAX(ℓiop)\mathbf{p}_i^{\text{op}} = \text{SOFTMAX}(\ell_i^{\text{op}})piop=SOFTMAX(ℓiop) 得到一个三维概率向量,例如 [P(skip),P(keep),P(repeat)][P(\text{skip}), P(\text{keep}), P(\text{repeat})][P(skip),P(keep),P(repeat)]。如果结果是 [0.05,0.90,0.05][0.05, 0.90, 0.05][0.05,0.90,0.05],表示模型强烈建议对该分段执行“keep”操作。
- 分段任务通过边界指示符上的二元交叉熵进行监督:
Lseg=−∑i=0D−1[ziseg∗logpiseg+(1−ziseg∗)log(1−piseg)].\mathcal{L}_{\text{seg}} = -\sum_{i=0}^{D-1} \left[ \mathbf{z}_i^{\text{seg}*} \log p_i^{\text{seg}} + \left(1 - \mathbf{z}_i^{\text{seg}*}\right) \log\left(1 - p_i^{\text{seg}}\right) \right].Lseg=−i=0∑D−1[ziseg∗logpiseg+(1−ziseg∗)log(1−piseg)].
ziseg∗z_i^{\text{seg}*}ziseg∗是否=1,即真实标签是“是不是起点”。如果真实标签是1(是起点),而预测概率 pisegp_i^{\text{seg}}piseg 不接近1,接近0,损失(惩罚)就会很大。
- 操作预测采用掩码交叉熵,损失仅在段起始位置处计算。设置掩码 mi=ziseg∗m_i = \mathbf{z}_i^{\text{seg}*}mi=ziseg∗,损失计算如下:
Lop=−∑i=0D−1mi⋅logpiop[ziop∗].\mathcal{L}_{\text{op}} = -\sum_{i=0}^{D-1} m_i \cdot \log \mathbf{p}_i^{\text{op}}\left[\mathbf{z}_i^{\text{op}*}\right].Lop=−i=0∑D−1mi⋅logpiop[ziop∗].
掩码的作用:操作标签只在分段起点有效,只有当第i层是真实分段起点时,mi=1m_i=1mi=1,损失函数才计入第i层的计算;否则 mi=0m_i=0mi=0,损失为0。
- 最终的目标函数为
L=Lseg+Lop\mathcal{L} = \mathcal{L}_{\text{seg}} + \mathcal{L}_{\text{op}}L=Lseg+Lop
通过这两个损失函数的结合,模型既学会了如何“切分”模型层(LsegL_{\text{seg}}Lseg),也学会了对切分出来的段块执行什么操作(LopL_{\text{op}}Lop)。
-
推理时的程序解码:在推理期间,执行程序分两阶段解码。
- 首先,通过对分割预测ℓsegℓ^{seg}ℓseg设定阈值,以确定分段边界。如果分段超出最大长度限制KmaxK_{max}Kmax,则插入额外边界以满足该限制,由此得到分段起始位置{sj}\{s_j\}{sj}。
- 基于已确定的分段,计算每个分段起点的操作对数概率:
logp(oj∣x,sj)=logSOFTMAX(ℓsjop)[oj]log p(o_j|x,s_j)= log SOFTMAX(\ell_{s_j}^{op})[o_j]logp(oj∣x,sj)=logSOFTMAX(ℓsjop)[oj]
没有通过局部贪心最大值法来独立地选取操作,而是考虑到片段之间存在非局部交互,采用小规模束搜索,保证产生全局更优、更一致的程序。束搜索在受限的空间内生成一组排序后的候选执行程序π(x)。最终,将每个候选程序按前面的“程序表示”,映射成一个具体的执行程序。
四、实验
4.1 实验设置
- 模型:LLaMA-3.2-3B-Instruct, Qwen1.5-MoE-A2.7BChat, Qwen2.5-3B-Instruct, Qwen3-8B.
- 数据集:DART-Math(分布内基准测试,按难度划分的训练/测试集),ASDiv和MAWPS、MMLU-Pro(分布外基准测试,基于全部难度等级的 DART-Math 训练数据集的并集进行训练,并采用零样本方式开展评估。)
- 指标:pass@k
- 基线方法:
- Base(τ = 0)采用温度参数τ = 0的贪心解码。
τ 代表温度参数。当 τ = 0时,模型在生成每一个词时,总是选择概率最高的那个词。
- Base(采样)利用随机解码,选取τ ∈ {0.3, 0.7, 1.0}生成k个输出样本,并选取不同温度下的最优结果,在不改动模型内部执行流程的前提下提升输出多样性。
- DR.LLM从执行路径中学习层路由策略,并在推理阶段启用该策略。
- ShortGPT依据预估重要性对网络层进行静态剪枝,得到浅层模型。
- MindSkip与FlexiDepth学习基于路由网络的动态深度策略,核心目标是优化推理效率。
- 其他:Mixture-of-Depths、LaCo、Mixture-of-Recursions等多种方案需要大量额外训练或是调整模型架构。
4.2 实验结果
-
精度提升:POLAR在Pass@1和Pass@k上均显著优于Base和现有动态深度方法。
- 分析:精度提升来源于模型内部潜在执行机制的优化,而非简单的输出多样性。

- 分析:精度提升来源于模型内部潜在执行机制的优化,而非简单的输出多样性。
-
效率提升:POLAR生成的程序往往更短,减少了平均执行层数,降低了端到端延迟。
- 开销:预测网络带来的额外计算开销极低(<1%),远小于节省的推理时间。
-
测试时缩放:增加候选程序数量kkk,POLAR的性能单调提升,证明了执行程序空间探索的有效性。

-
跨领域泛化:在数学数据集上训练的POLAR,能成功泛化到MMLU-Pro、ASDiv等分布外数据集,证明其学到的是通用的计算控制策略。

学习方法:POLAR&spm=1001.2101.3001.5002&articleId=163282843&d=1&t=3&u=e53c2fe6036a4f2fb05eb19cffeb0834)
83

被折叠的 条评论
为什么被折叠?



