FlanT5-CoT-Specialization
从 GPT-3 蒸馏思维链推理能力到 FlanT5 小模型的开源实现,ICML 2023 论文配套代码
加载项目详情…
本应用为开源项目,仅供学习研究,请遵守其开源协议。
从 GPT-3 蒸馏思维链推理能力到 FlanT5 小模型的开源实现,ICML 2023 论文配套代码
加载项目详情…
本应用为开源项目,仅供学习研究,请遵守其开源协议。
想象一下这个场景:你让一个大语言模型做数学题,它要么直接给答案,要么"灵光一闪"冒出推理步骤,但答案却错了。更糟糕的是,你没法让它只展示正确的推理过程,因为它根本不知道"自己懂推理"和"自己会推理"是两回事。
这是 2023 年初大模型领域的一个真实痛点:如何让小参数模型也具备多步推理能力?来自复旦大学、苏格兰大学和 Allen Institute for AI 的研究团队在 ICML 2023 上发表的论文《Specializing Smaller Language Models towards Multi-Step Reasoning》,给出了一个令人眼前一亮的答案——通过从 GPT-3 的强大"老师" code-davinci-002 蒸馏 Chain-of-Thought(思维链)推理能力到 FlanT5 小模型。

图1:项目作者 FranxYao(姚飞)主页
在 ChatGPT 引发的大模型浪潮中,有一个被反复验证的规律:模型的参数量越大,涌现能力(emergent abilities)越强。多步推理——即"先理解问题,再分步推导,最后给出答案"的能力——正是一种典型的涌现能力。GPT-3(175B 参数)在链式推理任务上表现优异,但将这样的能力蒸馏到仅有 780M 参数的 FlanT5-base 时,效果往往大幅退化。
研究团队的核心洞察是:小模型"推理差"的根本原因不是参数量不够,而是训练数据的格式不对。传统方法将问题-答案对作为训练数据,小模型只能学到"看到这个题,输出那个答案"的模式匹配。而 Chain-of-Thought 推理的核心在于中间推理步骤——这些步骤本身就是宝贵的监督信号。研究团队没有简单地将"问题和答案"喂给模型,而是精心构造了四种数据格式,分别激活模型的不同能力。
项目的核心创新不在于模型架构的改动,而在于数据格式的系统性设计。作者将预训练数据组织为四种格式,每种格式对应不同的推理能力激活:
in-context answer-only(上下文内仅答案格式) 提供了标准few-shot学习范式,让模型在相似题目的提示下直接输出答案,主要训练模型的模式匹配和快速推理能力。in-context chain-of-thought(上下文内思维链格式) 在 few-shot 提示中加入完整的分步推导过程,让模型同时学习"如何展开推理"和"如何得出正确答案"。zero-shot answer-only(零样本仅答案格式) 则训练模型在没有任何示例的情况下直接回答,考验模型的泛化能力。zero-shot chain-of-thought(零样本思维链格式) 是最具挑战性的设置——不提供任何示例,要求模型自发产生"首先...然后...最后..."的分步推理过程。
这四种格式的数据并非随意构造。作者通过 code-davinci-002 生成高质量的思维链数据,利用该模型强大的代码推理能力生成每一步的概率分布,再通过动态时间规整(Dynamic Time Warping,DTW)算法将 code-davinci-002 的 token 序列与 FlanT5 的 token 序列对齐,最终得到可供蒸馏的软标签数据。仓库中的 dev_align_codex_to_flan_t5_dtw.ipynb 和 processed_data/ 目录记录了这个复杂的对齐过程。
项目的代码结构非常清晰,体现了"复杂留给数据,简洁留给代码"的设计哲学。核心训练脚本 train_distill_simple.py 仅有几百行,使用 PyTorch Lightning 封装训练流程,支持分布式多卡训练。src/data_utils.py 包含完整的数据加载和预处理逻辑,处理从原始 .pkl 文件到 PyTorch DataLoader 的整个 pipeline。src/trainer_distill.py 定义了自定义的 DistillTrainer,继承自 Transformers 的 Seq2SeqTrainer。
仓库中的 prompt 库(lib_prompt/)存放了不同实验配置的提示模板,包括原始提示、简化提示、随机化提示等多种变体,用于对比实验。实验脚本 experiments.md 和 experiments_11b.md 详细记录了不同超参数配置下的训练命令,包括学习率(0.0005)、batch size(3B)、梯度累积步数(30)等关键参数。
值得注意的是,作者特别提到他们"没有时间实现 DeepSpeed/FairScale/PyTorch FSDP",但如果需要支持更大的模型 Wrapping DeepSpeed 应该"非常 straightforward"。这既是作者诚实的自我评价,也为社区贡献者留下了改进空间。
论文的核心实验在 GSM8K(小学数学)、BBH(Big Bench Hard)和 ASDIV 等推理基准上进行。结果显示,经过思维链特化训练的 FlanT5-11B 在 GSM8K 上的表现显著提升,蒸馏后的模型在保持较小参数量的情况下,推理能力逼近 GPT-3 175B 的水平。这一结果有力地证明了:通过精心设计的数据格式,可以在小模型上激活原本只在大模型上才涌现的推理能力。
项目还提供了多个 Jupyter Notebook 用于可视化实验结果:dev_process_codex_outputs.ipynb 可视化 Codex 模型的输出概率分布,dev_align_codex_to_flan_t5_dtw.ipynb 展示动态时间规整对齐效果,inspect_processed_data.ipynb 则是理解四种数据格式的最佳入口。
项目并非没有局限。首先,数据依赖严重:整个蒸馏流程依赖 Google Drive 分享的预训练数据(processed_data 目录),且 code-davinci-002 作为教师模型需要 OpenAI API 访问权限,这限制了复现的便利性。其次,训练稳定性:experiments.md 中明确标记了 model_version=0.0.2.2.1 为 BUGGY 版本,部分对比实验(如对比损失 contrastive loss)也标注为"not good",说明最佳超参数配置并非一蹴而就。第三,计算资源门槛:尽管是小模型,FlanT5-xl/11B 仍需要至少 24GB 显存(单卡 A100 或双卡 V100),个人开发者难以轻易复现。
FlanT5-CoT-Specialization 的价值不仅在于论文本身,更在于它开启了一个新的研究方向:如何通过数据工程而非模型 scaling 来激活小模型的能力上限。在此之后,Orca、WizardLM 等工作相继跟进,通过"模仿学习"和"课程学习"策略进一步探索了小模型专精化的可能性。
这一方向对工业界同样有重要意义:企业无需部署千亿参数的 GPT-4,通过对中等规模模型(如 FlanT5、LLaMA)进行任务特化训练,可以在特定场景(代码生成、数学推理、客服对话)上达到接近大模型的效果,同时大幅降低成本。
项目链接:https://github.com/FranxYao/FlanT5-CoT-Specialization
论文:arXiv:2301.12726
许可:MIT License
技术栈:Python / PyTorch / Transformers / Hydra Core / PyTorch Lightning / Jupyter Notebook