发表时间: 2024-02 · arXiv:2402.03300
原文: https://arxiv.org/abs/2402.03300
文章标题:DeepSeekMath: 推动开放语言模型中数学推理的极限
作者/机构:Zhihong Shao, Peiyi Wang, Qihao Zhu, Runxin Xu, Junxiao Song, Xiao Bi, Haowei Zhang, Mingchuan Zhang, Y.K. Li, Y. Wu, Daya Guo; 来自 DeepSeek-AI, 清华大学, 北京大学
一句话结论 本文通过构建1200亿 tokens 的高质量数学语料库并提出无需评论家模型的 GRPO 强化学习算法,训练出了 DeepSeekMath 7B 模型,在数学推理任务上以小模型尺寸达到了与 540B 闭源模型相媲美的顶尖性能。
要解决什么问题 原有的大语言模型在提升数学推理能力时主要面临两个层面的卡点。在预训练阶段,高质量数学数据的规模和来源严重受限,以往模型高度依赖 arXiv 论文,但这类数据规模容易见顶,且对基础数学解题能力的实际提升并不明显。在强化学习微调阶段,主流的近端策略优化算法 PPO 存在严重的显存和计算资源瓶颈。PPO 在训练策略模型时,必须同时训练一个与策略模型参数量相当的价值函数(评论家模型)来估计基线以减少方差,导致显存消耗翻倍。此外,在语言模型生成任务中,通常只有序列末尾的最后一个 token 会被赋予分数,这种极其稀疏的奖励信号使得训练一个能在每个 token 上都给出准确价值预测的评论家模型变得异常困难且低效,严重拖慢了强化学习的迭代效率。
怎么做的 为了绕开上述卡点,本文在数据构建和强化学习算法上进行了核心设计。在数据层面,核心思路是通过迭代式的分类器从海量网页中挖掘高质量数学文本。方法从高质量网页集合 OpenWebMath 作为种子语料库出发,训练 fastText 分类器去 Common Crawl 中召回相似网页。为解决正样本多样性不足,设计了基于域名的迭代扩充机制:将高比例召回的域名标记为数学相关域,人工标注未被召回的 URL 加入种子库重新训练。经过四轮迭代,构建了包含1200亿 tokens 的 DeepSeekMath 语料库。在强化学习层面,提出了组相对策略优化算法 GRPO,核心思路是摒弃 PPO 中极耗资源的价值模型,通过对同个问题采样多个输出,利用组内相对得分估计基线。关键设计包含两种优势计算方式。对于结果监督,GRPO 为每个问题采样 $K$ 个输出,通过奖励模型打分后减去均值并除以标准差进行归一化,将所有 token 的优势值设为该归一化奖励:$$ \hat{A}_{t,k} = \frac{r_k - \text{mean}(r)}{\text{std}(r)} $$对于过程监督,则在每个推理步骤结束时赋予奖励,将 token 的优势值计算为后续所有步骤归一化奖励之和:$$ \hat{A}_{t,k} = \sum_{T_{\text{step}}(j) \ge t} R'_{\text{step}(j)_k} $$这种组内相对比较机制完美契合了奖励模型基于比较数据训练的本质,彻底免去了额外训练价值模型的显存开销。此外,模型在预训练前先进行代码训练,显著提升了后续的数学推理能力。
效果如何 实验基于 7B 参数规模的模型展开,预训练使用5000亿 tokens 混合数据,并在77.6万条数据上进行指令微调,最后仅用约14.4万条思维链数据进行强化学习。对比基线涵盖代表通用大模型的 Mistral 7B、Qwen 72B,代表数学专项预训练的 Llemma 34B,代表数学指令微调与强化学习的 WizardMath、ToRA 34B,以及代表闭源最高水平的 GPT-4、Gemini Ultra 和 Minerva 540B。量化结果显示,预训练阶段的 DeepSeekMath-Base 7B 在 MATH 数据集上比现有开源基础模型高出超10个绝对百分点,并超越了参数量大77倍的 Minerva 540B。强化学习后的 DeepSeekMath-RL 7B 在无工具思维链推理下,在 GSM8K 和 MATH 上分别达到 88.2% 和 51.7% 的准确率,击败所有 7B 到 70B 开源模型。在允许使用 Python 工具辅助的设置下,MATH 准确率接近 60%。实验也揭示了局限:单纯增加 arXiv 论文数据并未带来显著性能提升;若在单阶段预训练中强行混合代码和数学数据,受限于小模型容量,反而会损害无工具的纯数学推理能力。
本文介绍了 DeepSeekMath,一个在数学推理方面取得显著性能的领域特定语言模型。其核心贡献可分为可扩展的数学预训练和对强化学习的探索与分析两个方面。
在可扩展的数学预训练方面:
1. 构建大规模高质量数学语料库:通过精心设计的数据筛选流程,从公开的 Common Crawl 数据中成功构建了包含1200亿词元(token)的高质量数学语料库——DeepSeekMath Corpus。该语料库的规模远超之前的研究,例如是 Minerva 使用的数学网页数据的近7倍,是 OpenWebMath 的9倍。
2. 验证小模型+高质量数据的有效性:预训练的 DeepSeekMath-Base 7B 模型性能可与比其大77倍的 Minerva 540B 相媲美,证明了模型参数量并非数学推理能力的唯一关键因素,在高质量数据上训练的小模型同样能取得强大性能。
3. 探索代码训练对数学推理的助益:研究发现,在进行数学训练之前先进行代码训练,能够提升模型在带工具和不带工具使用场景下的数学问题解决能力,为“代码训练能否提升推理能力”这一长期问题提供了肯定的部分答案。
4. 对 arXiv 论文数据的反思:实验表明,尽管在许多数学相关论文中普遍使用 arXiv 论文进行训练,但在本文所采用的数学基准测试中,并未带来显著的性能提升。
在对强化学习的探索与分析方面:
1. 提出高效的强化学习算法 GRPO:引入了组相对策略优化(Group Relative Policy Optimization, GRPO),这是一种 PPO 的变体。GRPO 无需评论家(critic)模型,而是通过组内分数来估计基线,从而显著减少了与 PPO 相比的训练资源消耗。
2. 验证 GRPO 的有效性:仅使用指令微调数据,GRPO 就显著提升了 DeepSeekMath-Instruct 模型的性能,并且在强化学习过程中观察到模型在领域外任务上的性能也有所提升。
3. 提供统一的分析范式:提出了一个统一的范式来理解包括 RFT、DPO、PPO 和 GRPO 在内的不同训练方法。通过广泛实验深入研究了该范式的关键要素(如在线与离线训练、结果与过程监督、单轮与迭代强化学习等),并基于此探讨了强化学习有效的原因,为未来实现更有效的强化学习总结了潜在方向。
构建DeepSeekMath语料库的迭代流程:本文概述了从 Common Crawl 构建 DeepSeekMath 语料库的过程。如图2所示,这是一个迭代流程,展示了如何从一个种子语料库(例如,一个小型但高质量的数学相关数据集)开始,系统地从 Common Crawl 中收集大规模数学语料库。该方法同样适用于编码等其他领域。
初始数据收集与分类器训练:首先,选择高质量数学网络文本集合 OpenWebMath 【【35】,OpenWebMath: An open dataset of high-quality mathematical web text,2023】作为初始种子语料库。利用这个语料库,训练一个 fastText 模型【【22】,Fasttext.zip: Compressing text classification models,2016】来召回更多类似 OpenWebMath 的数学网页。具体地,从种子语料库中随机选择50万个数据点作为正向训练样本,并从 Common Crawl 中选择另外50万个网页作为负向样本。使用一个开源库进行训练,配置向量维度为256,学习率为0.1,词n-gram的最大长度为3,词出现的最小次数为3,训练轮数为3。为了减小原始 Common Crawl 的规模,采用了基于URL的去重和近去重技术,最终得到400亿个HTML网页。然后,使用 fastText 模型从去重后的 Common Crawl 中召回数学网页。为了过滤掉低质量的数学内容,根据 fastText 模型预测的分数对收集到的页面进行排序,并只保留排名靠前的页面。保留的数据量通过对前40B、80B、120B和160B词元进行预训练实验来评估。在第一次迭代中,选择保留前40B个词元。
迭代优化与语料库扩充:在第一次数据收集迭代后,由于 fastText 模型训练所用的正样本集多样性不足,许多数学网页仍未被收集。因此,通过识别更多的数学网页来源来丰富种子语料库,从而优化 fastText 模型。具体操作是,首先将整个 Common Crawl 组织成不相交的域(domain),一个域定义为共享相同基础URL的网页。对于每个域,计算在第一次迭代中被收集的网页百分比。网页收集比例超过10%的域被归类为数学相关域(例如,mathoverflow.net)。随后,在这些已识别的域中手动标注与数学内容相关的URL(例如,mathoverflow.net/questions)。与这些URL相关但未被收集的网页将被添加到种子语料库中。这种方法能够收集更多的正样本,从而训练出更优的 fastText 模型,在后续迭代中召回更多的数学数据。经过四轮数据收集迭代,最终得到3550万个数学网页,总计1200亿词元。在第四次迭代中,注意到近98%的数据已在第三次迭代中被收集,因此决定停止数据收集。
基准测试污染清理:为避免基准测试污染,遵循 Guo 等人【【15】,Deepseek-coder: When the large language model meets programming – the rise of code intelligence,2024】的方法,过滤掉包含英文数学基准(如GSM8K【【9】,Training verifiers to solve math word problems,2021】和MATH【【17】,Measuring mathematical problem solving with the math dataset,2021】)以及中文基准(如CMATH【【57】,Cmath: Can your language model pass chinese elementary school math test?,2023】和AGIEval【【64】,AGIEval: A human-centric benchmark for evaluating foundation models,2023】)中的问题或答案的网页。过滤标准如下:任何包含与评估基准中任意子字符串完全匹配的10-gram字符串的文本段都将从数学训练语料库中移除。对于长度小于10-gram但至少有3-gram的基准文本,采用精确匹配来过滤受污染的网页。
优化器与学习率:遵循DeepSeek LLM的训练实践,使用AdamW优化器【【28】,Decoupled weight decay regularization,2017】,其中 $\beta_1 = 0.9$, $\beta_2 = 0.95$, weight_decay = 0.1。采用多步学习率调度,学习率在2000个预热步骤后达到峰值,在训练过程的80%后降至峰值的31.6%,在90%后进一步降至峰值的10.0%。学习率最大值设为5.3e-4,批处理大小为400万词元,上下文长度为4K。
高质量:使用少样本思维链提示【【56】,Chain-of-thought prompting elicits reasoning in large language models,2022】在8个数学基准上评估下游性能。如表1所示,在DeepSeekMath语料库上训练的模型性能明显领先。图3显示,在处理了500亿词元(相当于Proof-Pile-2的1个完整epoch)时,该模型表现优于在Proof-Pile-2上训练的模型,表明DeepSeekMath语料库的平均质量更高。
多语言性:DeepSeekMath语料库包含多种语言的数据,主要以英语和中文为代表。如表1所示,使用DeepSeekMath语料库进行训练,可以同时提升英语和中文的数学推理性能。相比之下,现有的以英语为中心的数学语料库在中文数学推理上的提升有限,甚至可能产生负面影响。
逐步推理评估结果:如表2所示,DeepSeekMath-Base 7B在所有八个基准上均领先于开源基础模型(包括广泛使用的通用模型Mistral 7B【【21】,Mistral 7b,2023】和近期发布的、在Proof-Pile-2上进行数学训练的Llemma 34B【【3】,Llemma: An open language model for mathematics,2023】)。值得注意的是,在竞赛级别的MATH数据集上,DeepSeekMath-Base比现有开源基础模型高出超过10个绝对百分点,并超过了规模大77倍的闭源基础模型Minerva 540B【【25】,Solving quantitative reasoning problems with language models,2022a】,后者基于PaLM【【26】,Solving quantitative reasoning problems with language models,2022b】并在数学文本上进行了进一步训练。
使用工具的数学问题解决:在GSM8K和MATH上,使用少样本思路编程提示(program-of-thought prompting)【【8】,Program of thoughts prompting: Disentangling computation from reasoning for numerical reasoning tasks,2022;【13】,PAL: programaided language models,2023】评估了程序辅助的数学推理能力。模型被提示通过编写Python程序来解决每个问题,其中可以利用math和sympy等库进行复杂计算。程序的执行结果被评估为答案。如表3所示,DeepSeekMath-Base 7B的性能超过了之前的最先进模型Llemma 34B。
形式化数学:形式化证明自动化有助于确保数学证明的准确性和可靠性,并提高效率,近年来受到越来越多的关注。我们在非正式到正式的证明任务【【20】,Draft, sketch, and prove: Guiding formal theorem provers with informal proofs,2022】上评估了DeepSeekMath-Base 7B,该任务是根据一个非正式陈述、该陈述的形式化对应物以及一个非正式证明来生成一个形式化证明。我们在miniF2F【【63】,Minif2f: a cross-system benchmark for formal olympiad-level mathematics,2021】(一个奥林匹克级别的形式化数学基准)上进行评估,并使用少样本提示为每个问题生成Isabelle中的形式化证明。遵循Jiang等人【【20】,Draft, sketch, and prove: Guiding formal theorem provers with informal proofs,2022】的方法,我们利用模型生成证明草图,并执行现成的自动证明器Sledgehammer【【36】,Three years of experience with sledgehammer, a practical link between automatic and interactive theorem provers,2010】来填补缺失的细节。如表3所示,DeepSeekMath-Base 7B在证明自动形式化方面表现出强大的性能。
自然语言理解、推理和代码能力:在MMLU【【16】,Measuring massive multitask language understanding,22】上评估模型的自然语言理解能力,在BBH【【46】,Challenging big-bench tasks and whether chain-of-thought can solve them,2022】上评估推理能力,在HumanEval【【7】,Evaluating large language models trained on code,2021】和MBPP【【2】,Program synthesis with large language models,2021】上评估编码能力。如表4所示,DeepSeekMath-Base 7B在其前身DeepSeek-Coder-Base-v1.5【【15】,Deepseek-coder: When the large language model meets programming – the rise of code intelligence,2024】的基础上,在MMLU和BBH上的性能有显著提升,表明数学训练对语言理解和推理有积极影响。此外,通过在持续训练中加入代码词元,DeepSeekMath-Base 7B有效地保持了DeepSeek-Coder-Base-v1.5在两个编码基准上的性能。总体而言,DeepSeekMath-Base 7B在三个推理和编码基准上显著优于通用模型Mistral 7B【【21】,Mistral 7b,2023】。
PPO算法回顾:近端策略优化(PPO)【【42】,Proximal policy optimization algorithms,2017】是一种在LLM的RL微调阶段广泛使用的演员-评论家(actor-critic)RL算法【【34】,Training language models to follow instructions with human feedback,2022】。它通过最大化以下代理目标来优化LLM:
其中 $\pi_{\theta}$ 和 $\pi_{\text{old}}$ 分别是当前和旧的策略模型,$x, y$ 分别是从问题数据集和旧策略 $\pi_{\text{old}}$ 中采样的问题和输出。$\epsilon$ 是PPO中为稳定训练引入的与裁剪相关的超参数。$A_{t}$ 是优势(advantage),通过广义优势估计(GAE)【【41】,High-dimensional continuous control using generalized advantage estimation,2015】计算得出,基于奖励 {$r_{\ge t}$} 和一个学习到的价值函数 $V_{\phi}$。因此,在PPO中,需要与策略模型一同训练一个价值函数,并且为了减轻对奖励模型的过度优化,标准方法是在每个词元(token)的奖励中加入一个来自参考模型的逐词元KL惩罚项【【34】,Training language models to follow instructions with human feedback,2022】,即:
其中 $r_{\psi}$ 是奖励模型,$\pi_{\text{ref}}$ 是参考模型,通常是初始的SFT模型,$\beta$ 是KL惩罚项的系数。
GRPO算法提出:由于PPO中使用的价值函数通常是与策略模型规模相当的另一个模型,这带来了巨大的内存和计算负担。此外,在RL训练中,价值函数在计算优势时作为基线以减少方差。但在LLM场景下,通常只有最后一个词元被奖励模型赋予奖励分数,这可能使训练一个在每个词元上都准确的价值函数变得复杂。为了解决这个问题,如图4所示,我们提出了组相对策略优化(GRPO),它避免了PPO中对额外价值函数近似的需求,而是使用对同一问题采样的多个输出的平均奖励作为基线。更具体地,对于每个问题 $x$,GRPO从旧策略 $\pi_{\text{old}}$ 中采样一组输出 {$y_1, y_2, \cdots, y_K$},然后通过最大化以下目标来优化策略模型:
其中 $\beta$ 和 $\gamma$ 是超参数,$\hat{A}_{t,k}$ 是仅基于组内输出的相对奖励计算出的优势,具体将在后续小节中详述。GRPO利用组相对方式计算优势,与奖励模型的比较性质非常吻合,因为奖励模型通常是在对同一问题的输出进行比较的数据集上训练的。还需注意,GRPO不是在奖励中添加KL惩罚项,而是通过直接将训练策略与参考策略之间的KL散度添加到损失中来进行正则化,避免了 $\hat{A}_{t,k}$ 计算的复杂化。
GRPO算法流程:
输入:初始策略模型 $\pi_{\theta}^{\text{init}}$;奖励模型 $r_{\psi}$;任务提示 D;超参数 $\beta, \gamma, K$
1: 策略模型 $\pi_{\theta} \leftarrow \pi_{\theta}^{\text{init}}$
2: for 迭代 = 1, . . . , I do
3: 参考模型 $\pi_{\text{ref}} \leftarrow \pi_{\theta}$
4: for 步骤 = 1, . . . , M do
5: 从 D 中采样一批 $D_b$
6: 更新旧策略模型 $\pi_{\text{old}} \leftarrow \pi_{\theta}$
7: 对每个问题 $x \in D_b$,采样 $K$ 个输出 {$y_k$}$_{k=1}^{K} \sim \pi_{\text{old}}(\cdot | x)$
8: 通过运行 $r_{\psi}$ 计算每个采样输出 $y_k$ 的奖励 {$r_k$}$_{k=1}^{K}$
9: 通过组相对优势估计计算 $y_k$ 的第 $t$ 个词元的 $\hat{A}_{t,k}$。
10: for GRPO 迭代 = 1, . . . , $N$ do
11: 通过最大化GRPO目标(公式21)更新策略模型 $\pi_{\theta}$
12: 通过使用回放机制的持续训练来更新 $r_{\psi}$。
输出:$\pi_{\theta}$
KL散度估计器:与公式(2)中使用的KL惩罚项不同,我们使用以下无偏估计器【【40】,Approximating kl divergence,2020】来估计KL散度:
这个估计器保证为正。
评估结果:表5展示了开源和闭源模型在有无工具辅助推理下,在中英文基准上的性能。我们发现:1) DeepSeekMath-RL 7B在使用思维链推理时,在GSM8K和MATH上分别达到了88.2%和51.7%的准确率。这一性能超过了所有7B到70B范围内的开源模型以及大多数闭源模型。2) 关键的是,DeepSeekMath-RL 7B是从DeepSeekMath-Instruct 7B开始,仅在GSM8K和MATH的思维链格式指令微调数据上进行训练。尽管其训练数据范围有限,它在所有评估指标上都优于DeepSeekMath-Instruct 7B,展示了强化学习的有效性。
本节分享了在预训练和强化学习实验中的发现。
实验设置:为了研究代码训练如何影响数学推理,我们试验了以下两阶段训练和单阶段训练设置:
实验结果:表6和表7展示了不同训练设置下的下游性能。
实验结果:在我们的实验中,我们分别在每个ArXiv语料库上训练DeepSeek-LLM 1.3B 1500亿词元和DeepSeek-Coder-Base-v1.5 7B 400亿词元。结果表明,ArXiv论文在提升数学推理方面似乎无效。当仅在ArXiv语料库上训练时,两个模型在本研究采用的各种复杂度的数学基准上均未显示显著提升,甚至出现性能下降。这些基准包括定量推理数据集如GSM8K和MATH(表8),多项选择挑战如MMLU-STEM(表8),以及形式化数学如miniF2F(表9)。
结论局限性:然而,这个结论有其局限性,应谨慎对待。我们尚未研究: