HyperAIHyperAI

Command Palette

Search for a command to run...

从预训练到后训练理解推理能力

Jingyan Shen Ang Li Salman Rahman Yifan Sun Micah Goldblum Matus Telgarsky Pavel Izmailov

摘要

强化学习(RL)已成为提升大语言模型(LLMs)在复杂推理任务上表现的核心手段,然而RL后训练大多被孤立地研究,与其之前的预训练相分离。因此,两个基本问题仍未得到解答:(1)预训练选择(模型规模、数据)如何影响RL计算投入的回报,以及(2)RL究竟对模型做了什么?这些问题在标准LLM设定中难以研究:预训练语料库庞大且不受控,使得难以将行为归因于预训练或RL,而跨两个阶段的系统性计算扫描成本过高。为应对这些挑战,我们使用国际象棋作为受控测试平台,研究从预训练到后训练全流程中的推理能力。我们遵循标准LLM训练流程,在人类棋局上预训练规模从5M到1B参数的语言模型,在合成推理轨迹上进行监督微调,并在可验证奖励的象棋谜题上运行RL。利用这一框架,我们建立了一条连接预训练与RL的缩放定律:在给定RL计算水平下的后RL表现可由预训练损失很好地预测,且RL奖励曲线的斜率随预训练令牌数近似线性提升。除缩放外,我们发现RL并非简单地锐化SFT策略:在简单谜题上,它放大SFT策略已偏好的正确走法;而在困难谜题上,它浮现出SFT下几乎不存在的正确走法。我们进一步在数学领域训练一个1B语言模型以检验象棋之外的发现,结果出现了相同的预测模式:预训练更久的检查点达到更高的后RL表现,并在RL下提升更快。总之,我们提供了预训练与RL接口的定量描述,以及一个用于研究从预训练到后训练全流程推理科学的受控测试平台。

一句话总结

纽约大学、Modal Labs等机构的研究人员将国际象棋作为受控测试平台,研究从预训练到强化学习的完整流程,建立了一条缩放定律:强化学习后的性能可由预训练损失预测,且强化学习奖励提升斜率随预训练token数量线性增长;同时揭示出,在简单谜题上强化学习会放大正确走法,而在困难谜题上则会浮现出新的正确走法,这些发现也延伸到了数学领域。

核心贡献

  • 一个受控的国际象棋测试平台复现了从预训练到强化学习后训练的完整大语言模型流程,能够系统性地改变预训练选择和强化学习计算量。
  • 联合缩放定律表明,强化学习后的性能水平可由预训练损失很好地预测,且强化学习奖励提升的斜率与预训练数据token数量近似呈线性缩放关系。
  • 机制分析揭示,在简单谜题上,强化学习放大了监督微调策略已经偏好的正确走法;而在困难谜题上,它浮现出原本几乎不存在的正确走法,但有时也会强化错误走法,这解释了为何pass@1的提升并不总能传递到pass@16。

引言

标准的大语言模型训练流程将大规模预训练与基于可验证奖励的强化学习后训练相结合,但固定计算预算在这两个阶段之间的最优分配仍不清楚。先前的工作分别针对预训练和强化学习后训练提出了缩放定律,然而尚无定量描述预训练与强化学习如何相互作用,且存在争议:强化学习仅仅是锐化已有行为,还是组合出新技能?作者通过构建一个受控的国际象棋测试平台来弥补这些空白,该平台模拟了大语言模型训练流程,能够对预训练和强化学习计算量进行系统性扫描。他们推导出一个联合缩放定律,表明预训练损失可预测强化学习后的性能,且强化学习的提升斜率随预训练数据规模增长;他们还分析了强化学习如何在不同难度的问题上差异化地重塑策略。

数据集

作者从Lichess平台构建了三个国际象棋数据集,所有数据集在棋盘局面层面互不相交,以防止数据污染。

  • 预训练语料库 来源:2022年在Lichess上进行的快棋和超快棋对局。 规模:540亿token。 用途:缩放扫描从此语料库中抽取不同数量的token预算来预训练基础模型。

  • 后训练谜题集 来源:15.6万道经过质量筛选的Lichess谜题。 结构:谜题被分为五个难度区间(B1至B5,从易到难)。 筛选:经过质量筛选,但正文未详述具体标准。 用途:用于预训练之后的后训练(监督微调)。

  • 评估基准 来源:1480道战术谜题,取自相同的来源和难度区间。 整理:在主题多样性和解法长度上进行了平衡。 报告:由于模型很少能解决B5区间的谜题,汇总的pass@k结果基于B1至B4计算;B5保留用于按难度分层的机制分析。

预训练数据用于语言建模,后训练谜题用于结合构建的推理轨迹进行指令微调,评估集用于衡量不同难度级别下的谜题解决准确率。

方法

作者利用国际象棋作为受控测试平台研究推理能力,设计了一个模拟标准语言模型范式的训练流程。整体框架包含三个顺序阶段:在人类对局数据上预训练、使用合成推理轨迹进行监督微调,以及在可验证谜题环境中进行强化学习。如下图所示,该流程系统性地将模型从基本的走法预测过渡到复杂的多步规划。

国际象棋表示与预训练 为处理国际象棋对局,作者将每局棋表示为玩家走法的交替序列,并序列化为离散的token。基于标准国际象棋记谱法,每个走法使用四token结构编码:piece\langle piece \ranglepiecesource\langle source \ranglesourcedestination\langle destination \rangledestinationflag\langle flag \rangleflag,其中flag token表示诸如王车易位或将杀等特殊动作。这种表示方式产生了一个紧凑的词表,大小为V=81|V|=81V=81。在预训练阶段,使用标准的下一token预测目标在大规模人类对局轨迹语料库上训练自回归策略,使模型学习合理走法序列的分布。

使用合成推理轨迹的监督微调 在后训练阶段,模型被训练来解决国际象棋谜题,需在给出最终走法前生成自己的推理轨迹。谜题环境被构建为一个多步交互决策过程。给定初始棋盘状态s0s_0s0和真实解法线路,模型仅作为求解方。在每一步ttt,模型观察当前状态sts_tst并提议一个候选走法ata_tat。环境根据真实解法验证ata_tat。若走法不匹配,则立即终止该回合;若走法正确,环境会将对手的回应oto_tot追加到上下文中,并将状态转移至st+1s_{t+1}st+1

为了在不依赖外部搜索算法的情况下激发上下文推理,作者构建了合成推理轨迹。他们从预训练的提议策略中采样KKK条可能的对局延续,并按它们的公共前缀合并成一棵以s0s_0s0为根的树结构。然后,该树按深度优先顺序序列化,形成推理轨迹rrrr=Tτ~1sepτ~2sepτ~m/Tr = \langle T \rangle \tilde{\tau}_1 \langle sep \rangle \tilde{\tau}_2 \langle sep \rangle \cdots \tilde{\tau}_m \langle /T \rangler=Tτ~1sepτ~2sepτ~m/T 其中τ~i\tilde{\tau}_iτ~i表示从根到叶的路径,T\langle T \rangleT/T\langle /T \rangle/T是界定思考标签的标记。随后,模型在拼接序列w=(r,τ)w = (r, \tau^\star)w=(r,τ)上进行训练,其中τ\tau^\starτ是真实解法延续。在训练过程中,对手走法的token对应的损失被屏蔽,确保模型仅针对推理轨迹和自身走法进行优化。

基于可验证奖励的强化学习 在监督微调策略的基础上,作者应用强化学习在可验证谜题环境中进一步优化模型。他们采用严格的二元结果奖励函数: R(ζ,s0)=1[a1=a1,,aH=aH]R(\zeta, s_0) = \mathbf{1}[a_1 = a_1^\star, \ldots, a_H = a_H^\star]R(ζ,s0)=1[a1=a1,,aH=aH] 其中ζ\zetaζ是包含推理轨迹和已执行走法序列的完整轨迹,(a1,,aH)(a_1^\star, \ldots, a_H^\star)(a1,,aH)是真实解法线路。仅当每一个执行的走法都与真实解法完全匹配时,模型才能获得1的奖励,这意味着只要出现一次错误便无奖励。策略使用分组相对策略优化(Group Relative Policy Optimization)进行优化,以提升其战略规划和走法选择能力。

实验

本研究以国际象棋谜题为推理测试平台,分析预训练和强化学习(RL)在从500万到10亿参数规模的模型上如何相互作用。关键发现是一条联合缩放定律:强化学习后的性能由预训练损失预测,而强化学习计算带来的提升速率取决于预训练token的数量,导致随着总预算增长,最优计算分配向强化学习倾斜。机制分析揭示,强化学习通过放大正确模式并发现部分尾部走法来重塑走法策略,但在更困难的任务上也会强化错误模式,从而限制了pass@k的提升;推理轨迹的改进主要来自更广泛的搜索而非更深入的规划。对数学推理的定性迁移显示出类似的模式,表明这些缩放关系可能超越国际象棋领域。


用 AI 构建 AI

从创意到上线——通过免费 AI 协同编码、开箱即用的环境和最优惠的 GPU 价格,加速您的 AI 开发。

AI 协同编码
开箱即用的 GPU
最优定价

HyperAI Newsletters

订阅我们的最新资讯
我们会在北京时间 每周一的上午九点 向您的邮箱投递本周内的最新更新
邮件发送服务由 MailChimp 提供