ARTICLE DETAIL

资讯详情

深耕网站建设与运营推广的一线实战洞察。

单卡A100 8小时训练小模型:SFT+GRPO实现循环思考与自我修正

单卡A100 8小时训练小模型:SFT+GRPO实现循环思考与自我修正 1. 项目定位一张A100、8小时从零训出会“循环思考”的小模型最近我做了一件挺有成就感的事只用一张A100花了8小时训练出一个会在回答问题之前先“想一想”、回答之后再“检查一遍”、发现错误还能“自己改过来”的小模型。说“从零训练”不是从随机权重预训练一个模型而是从一个通用的预训练语言模型底座出发针对“循环思考”这个专项能力做完整的增量训练。它一开始只会普通的“提问—回答”训练完以后它会生成“思考—初步回答—反思—修正”这样一条完整轨迹像是人类打草稿、检查、改错的过程。这个项目做完以后我把整个方案、数据格式、训练参数、踩过的坑都整理了一遍。如果你手里也有一张A100、想低成本验证强化学习的训练链路或者想让一个小模型具备自我修正能力这篇记录可以直接当作“抄作业”的参考。1.1 “循环思考”到底是什么为什么小模型需要它先说一个大家平时能观察到的现象小模型回答问题容易“嘴巴比脑子快”。你问它一道两步计算的数学题它经常跳步或者直接输出一个凭直觉得到的答案错了也不会纠正。而大模型在推理时通过内部的长思维链能大幅降低这类错误——但大模型的成本也摆在那里。“循环思考”这个能力说白了就是给模型外挂一套显式的思考流程先在心里推演给出一个初步答案然后跳出这个答案重新审视一遍发现漏洞以后修正它。这样做最大的好处是错误的答案不直接在最终回复中出现而是被“拦截”在反思环节。我一开始也疑惑小模型真的能学会这套流程吗后来实测发现它的学习能力远比我们想象中强。模型不需要理解复杂的“自我意识”只需要从训练数据里学会一种说话和作答的结构先不急着下结论而是把推导过程写出来写完之后回头检查如果发现问题就重新组织答案。这个能力本质上是一个序列到序列的格式学习问题小模型足够handle。1.2 为什么“一张A100 8小时”能完成训练先拆一下账。本来很多人一听“训练模型”下意识觉得要几十张卡跑一周。但这个项目能压缩到8小时核心在于三层第一选小底座。我用的是3B级别的开源模型而不是7B或者70B。3B模型在A100上的显存占用和计算量都很友好不管是全参微调还是后期RL阶段单卡都能扛住。如果你用1.5B时间还能更短。第二不是从头预训练而是增量训练。预训练需要看几十亿token而增量训练只需要看几万条高质量样本。这就好比一个是让一个婴儿从头认识世界另一个是让一个已经大学毕业的年轻人培训班。我们要做的是后者只教它“循环思考”这一种工作方式。第三训练框架选对了。这次全程用了常见的训练框架组合SFT阶段用SFTTrainerWrapper做监督微调RL阶段用GRPOTrainerWrapper做GRPO训练。这两个wrapper把数据加载、损失计算、路径采样、奖励收集全部封好了我不需要自己写训练循环省下了大把调试时间。我把三种方案的投入对比放在下面方便你理解为什么选这条路。方案训练时长硬件需求最终效果适合场景从零预训练3B模型数千GPU小时多卡集群通用能力弱还需额外微调科研探索不接地气通用底座全量微调数小时单卡A100足够目标能力突出通用性略降专项能力训练通用底座LoRA增量训练数小时单卡A100或更低目标能力可接受显存更省低成本快速验证1.3 技术选型底座模型、训练框架与训练策略底座我选的是Qwen系列3B指令版也就是Qwen2.5-3B-Instruct。选它的理由有几个首先是推理能力在3B梯队里算靠前的中文和英文都覆盖适合做中英文混合的评测其次是开源协议对二次训练友好不会在发布或商用环节卡你最后是生态成熟从HuggingFace的Transformers到TRL训练库都能直接加载。训练策略上我分了两个阶段。第一阶段用SFT监督微调目的是让模型快速学会“思考—初步回答—反思—修正”这个输出格式。这个阶段相当于给它看大量标准答案的范文让它照着格式写。第二阶段用GRPO一种强化学习训练策略相比PPO内存更省、对单卡更友好目的是用奖励函数告诉模型不是只要格式对就加分而是最终修正后的答案正确才加分。这样一来模型会从“背格式”进化到“真的想办法把答案改对”。很多人听到强化学习就头疼其实GRPO已经被封装得很好了。GRPOTrainerWrapper内部已经处理好了采样、批处理、KL散度约束这些细节你要做的只是定义一个奖励函数给它一个生成结果让它返回一个分数。我把奖励函数的设计放到后面的实操部分详细讲这里先不展开。2. 训练环境准备与数据构建2.1 A100环境下的显存规划与基础配置训练环境看起来简单但细节不少。我这里用的是A100 80GCUDA版本12.x配合PyTorch 2.x和Transformers 4.x系列。为什么强调A100因为80G显存在单卡场景下能给你充足的余量不仅模型能放得下还能同时负载训练时的激活值、优化器状态甚至RL阶段的rollout生成。照着我常用的配置来镜像nvcr.io/nvidia/pytorch:24.04-py3框架transformers、peft、trl、vllm精度bf16混合精度。A100虽然是老一代架构但对bf16有原生支持训练稳定性比fp16好很多后面不会突然出现loss变NaN。分布式单卡不涉及分布式但如果你有自己的训练习惯可以把DeepSpeed配置成zero-0没必要为3B模型上zero-2。说说显存的分配逻辑。3B模型的bf16权重大约6GB梯度再占6GBAdamW优化器状态一般用fp32保存一阶和二阶动量大约会占到24GB左右也就是说模型梯度优化器总共36GB左右。剩下的显存留给中间激活值、数据集token和RL阶段的生成缓冲区。这也是3B模型能在单卡A100上全参微调而7B只能退而求其次用LoRA的根本原因。如果你手里的卡不是80G而是40G甚至24G我也试过类似的方案把训练改成LoRArank16或者rank8效果依然能打只是模型可能少一点“灵性”。2.2 循环思考样本的结构设计与完整示例这个项目的关键之一不完全是“喂多少数据”而是“数据长什么样”。我构建样本的时候要求每一条训练数据都要体现“从粗到细、从错到对”的过程而且格式必须严格一致否则推理时模型不会稳定触发完整的循环思考。每条数据的对话历史我固定为四段式用特殊标记区分thought模型内部的推演过程answer初始答案reflection对初始答案的检查指出不足或错误revised修正后的最终答案下面给一个真实的训练样本这是从一个简单的比例题里截出来的我每周都会拿这种题目当校验集{ conversation: [ { role: user, content: 一个矩形的长增加20%宽减少20%它的面积变化了多少 }, { role: assistant, content: thought\n题目问面积变化百分比。我先设原长为L原宽为W原面积是L×W。变化后的长是1.2L变化后的宽是0.8W所以新面积 1.2L×0.8W 0.96LW。\n/thought\nanswer\n面积变化为减少4%。\n/answer\nreflection\n我重新核对一下计算1.2乘以0.8等于0.96确实比原来少0.04。这里有个常见的陷阱很多人会直接把20%和20%相减得到“不变”实际上因为是乘法运算结果不是0%。我的推导里考虑了这一点应该没问题。\n/reflection\nrevised\n面积减少4%。\n/revised } ] }注意这里的细节反思部分不能只是说“我做对了”那样模型会学成空洞的自夸。反思要引导模型真正“重算一遍”或“换一种思路验证”这样才能把正确率提上去。对于初始答案是错的样本反思部分要明确指出错处再在修正部分重新计算结果。数据总量上我准备了约2000条这样的完整轨迹其中约1500条是“初始答案错误、反思后修正正确”的样本剩下500条是“初始答案已经正确、但通过反思确认”的样本。这个配比非常重要如果全部是改错样本模型会倾向在第一步故意答错好让后面的修正显得有价值如果全部是正确的样本模型又学不会真正的纠错。混合配比才能让它既敢给出初步答案又愿意检查。2.3 数据预处理、验证集划分与防泄漏数据来源方面我没有全部依赖公开数据集而是采用“公开题目自己做答案轨迹 自造对抗样本”的组合。公开题目用数学、逻辑推理、常识判断这些答案容易客观校验的类型自造样本专门针对模型常见错误模式设计比如比例计算、否定句理解、单位换算等。预处理时要注意三点格式统一。所有标签都用完全一致的字符串不能一会儿用thought一会儿用[thought]推理时模型会认死理格式错一个字符都不会触发。去重。拿到公开数据集先跑一遍simhash或者简单的文本相似度把重复题去掉。否则训练集和验证集里可能出现几乎一样的题评估时会虚高。划分。按8:1:1切分成训练集、验证集、测试集。验证集用来做训练中间检查测试集留着最后统一评估。特别是自造的样本不能既出现在训练集又出现在测试集这种数据泄露会让你的正确率数字好看但一上真实场景立刻露馅。预处理完后把所有样本统一格式化成一个JSONL文件一行一个完整对话长度控制在1024个token以内超过的直接裁剪或丢弃。3B模型处理长文本的能力有限如果样本里面“思考”部分写太长反而会稀释训练信号。3. SFT GRPO 两阶段训练实操全流程3.1 第一阶段SFT让模型先把循环思考的“格式”跑通训练的第一步是用SFT把模型的输出格式掰过来。这个阶段比较简单数据喂进去模型学习“看到问题后应该按什么顺序说话”。我用的SFT配置是这样的模型Qwen2.5-3B-Instruct训练方式全参微调学习率5e-5使用cosine调度器前10%步数warmupbatch大小单卡batch设为4梯度累积8步等效batch为32训练轮数4个epoch最大序列长度1024精度bf16A100 80G跑全参3B很从容。2000条数据每条长度在几百token等效batch 324个epoch下来大概是250步左右实际训练时间在2到2.5小时之间。训练到一半的时候loss会降到1.0以下此时手动拿几条验证集数据测一下会看到模型已经能输出完整的thought、answer、reflection、revised结构只是修正后内容可能还是跟初始答案差不多的“假反思”。这一步我不建议为了省时间把LoRA拉出来。SFT阶段模型要学习新的输出格式参数更新范围越大格式学得越稳。全参在这个数据量下压力不大等到RL阶段再上LoRA更合理。3.2 第二阶段GRPO用奖励函数逼模型“真修正”SFT跑完以后格式是对了但内容质量还不行。要让模型真正学会“发现错误并修正”就得靠GRPO强化学习。先说GRPO比PPO在单卡场景下友好在哪。PPO需要同时维护Actor、Critic、Reference三个模型Critic本身就是很大的显存消耗GRPO去掉Critic用一组采样结果的相对优势来估计策略更新方向显存占用小得多训练难度也低。单卡A100做3B模型的RLGRPO是唯一实操可行的方案。奖励函数非常关键我把它拆成三个正交部分格式奖励输出必须包含四个标签且顺序正确得0.3分。答案正确奖励修正后的最终答案是否正确用规则匹配或一个简单的eval函数判断正确得1.0分。修正有效性奖励比较初始答案和修正后答案的对错变化。从错到对得0.5分从对到错扣1.0分都对或都错得0分。这里特别说一下第三个奖励。如果没有它模型很容易出现一个“投机行为”初始答案故意写得含糊或者错一点修正时再给正确答案这样既拿到“答案正确”的1分又拿到“从错到对”的0.5分。加了“从对到错扣1分”惩罚模型就不敢在原来已经答对的情况下乱改了。GRPO阶段我用LoRA做参数高效微调rank取16alpha取32target_modules覆盖所有linear层。学习率降到1e-6这比SFT阶段低了一个数量级。强化学习阶段的学习率必须保守否则几步更新就可能把底座模型在SFT阶段学到的格式能力冲掉。batch上每个prompt采样5个rollout结果等效batch 16训练约60步。这一步的实际训练时长在3到4小时之间取决于生成文本的长度。加上前面的SFT和数据处理整个项目控制在8小时内是妥当的。3.3 训练参数速查表与时间分配建议为了方便你对照复现我把关键参数和每个环节的耗时整理成一个表。阶段关键参数时长数据构建2000条四段式轨迹8:1:1划分1~1.5小时SFT全参5e-5batch 324 epochs2~2.5小时GRPOLoRA rank161e-660 stepskl_coef 0.053~4小时评估调试验证集抽查奖励函数调参0.5~1小时如果你的时间更紧张比如只有6小时我的建议是优先压缩数据构建时间从2000条减到1200条SFT训练降到3个epochGRPO阶段只跑30步。效果依然会比你未训练前好很多只是“从错到对”的稳定性会差一些。4. 效果验证与常见问题排障4.1 训练后模型真的“循环思考”了吗训练结束后我没有急着看各种花哨的指标而是先拿几个典型的测试题手动测。这里展示一个我在测试集里留的题目模型输出的片段已经过简写问题如果一个数除以3余2除以5余3这个数最小是多少 thought 我先找满足“除以5余3”的数3、8、13、18…… 再看哪个满足“除以3余2”。8除以3余2所以答案是8。 /thought answer这个数是8。/answer reflection 检查一下8÷3余2正确8÷5余3正确。题目要最小数8前面还有没有更小的3不满足所以8就是最小解。 /reflection revised最终答案是8。/revised注意看它已经学会在反思里进行“反向检验”而不是简单重复一遍答案。这就是GRPO奖励函数起的作用——模型知道“检查一下”会带来收益。我统计了测试集上60道题的正确率训练前的底座模型直接作答正确率约38%训练后的模型初始答案正确率约45%但经过反思修正后最终答案正确率到了71%。这个提升在3B小模型上已经非常明显了。这里有一个使用技巧推理时把温度设为0.7并可以在提示词里加一句“请先思考再回答最后检查修正”。如果温度太低比如0.1模型有时候会跳过思考直接输出最终答案因为贪心解码会更倾向于概率最高的短路径。4.2 训练中常见的4个坑与排查方法这个项目最花时间的往往不是训练本身而是排障。我把过程中遇到过的几个典型问题整理成表你遇到类似情况可以直接对照。现象可能原因解决建议训练中loss变成NaN学习率过大、数据里有异常超长样本减小学习率到3e-5以内启用梯度裁剪排查是否有超过1024 token的异常数据模型会输出反思但不会改答案反思和修正之间的数据映射不足增加“初始错误—反思—修正正确”样本的比例并在奖励函数中加大答案正确权重修正后反而把对的改错缺少对“对改错”的惩罚在修正有效性奖励中加负分项从对到错直接扣1分生成内容过长甚至无限续写样本平均长度太长模型把循环思考当成了无限续写限制max_new_tokens在512以内训练数据中混入普通单轮回答样本平衡其中“会反思但不会改答案”这个问题极其隐蔽。我一开始看到模型输出了正确的reflection标签就以为训练成功了后来把生成结果完整展开一看发现反思内容写的是“我的答案是对的不需要修改”然后修正部分跟初始答案一字不差。这就是典型的“格式学会了能力没学会”。解决办法只有一个狠下心调整数据配比和奖励系数逼着模型在反思时真的“发现问题”。4.3 提高训练效率的几个小技巧最后分享几个可以帮你少走弯路的经验。第一强化学习阶段一定要边训边测。不要等到几十个step跑完再看最终结果那样一旦奖励函数设计有问题浪费几个小时。我一般每5个step就在10条固定测试样本上跑一次手动评测肉眼扫一眼输出结构发现问题立刻停。第二反复调整奖励权重时每次只改一个参数。我试过一次同时调答案正确分和修正有效性分结果出问题后根本说不清是哪个权重改坏了。奖励系数调参要像做A/B测试一次只动一个变量。第三如果你后续想把方案迁移到7B或14B模型注意把RL阶段的LoRA换成全参会更稳。大模型的生成能力更强对GRPO的参数波动更敏感全参训练虽然慢但不容易出现策略崩掉。最后再分享一点实际体会整个项目走下来我最大的体会是训练一个会“循环思考”的小模型真正的瓶颈不是显卡而是你对数据和奖励函数的理解。8小时里真正跑训练可能只有6小时剩下的2小时都在反复看生成结果、挑数据问题。很多人以为有了A100就能一切自动化其实不然。如果你也想在类似方向上试水我建议从“3B模型 600条数据 GRPO 30步”的最小闭环开始先跑通流程再逐步加大数据量。这样即使翻车损失的时间也就一两个小时。我自己最开始就是先用最小闭环验证了奖励函数有效才敢放开手做完整训练。这个思路后续还可以扩展到更多场景比如让模型在写代码前先自检、在翻译后对照原文检查、在长文回答中分段推理。只要“初答—反思—修正”这个循环的奖励函数设计得当小模型都能学得有模有样。如果后面有了新的进展我再回来补充。
返回列表