从零训练小语言模型:预训练、CPT、SFT、PEFT、蒸馏与DPO全流程实战
1. 为什么我要从零折腾一个小语言模型先说结论从零训练一个小语言模型不是让你去跟那些千亿参数的大模型硬碰硬而是用一套完整的流程把预训练、CPT、SFT、PEFT、蒸馏、DPO这几个关键环节全部亲手跑一遍。这套流程走下来你对语言模型的理解会从“调API的”变成“知道模型肚子里发生了什么的”。我之所以选 Xihe 这个架构来做这件事理由很实在。Xihe 是一个相对轻量的中文语言模型结构参数量可控单卡甚至消费级显卡就能跑起来预训练的小规模实验。它的设计思路借鉴了 RoBERTa 中文预训练模型的一些经验同时保留了自回归生成的灵活性。对于想搞明白“预训练到底在训什么”的人来说拿它上手比直接去碰那些动辄几十G显存的模型要友好得多。这篇文章适合谁看如果你已经会调用现成的大模型API但对其中的训练流程一知半解或者你做过一些微调但从来没走过完整的“预训练到对齐”全链路再或者你是个刚入门的算法工程师想找一个能跑通全流程的小项目练手——那这篇内容就是写给你的。我会把每个环节的核心原理、实操步骤、参数选择依据、踩坑记录全部摊开讲代码和配置能给的都给让你看完能直接复现。整个流程的骨架是这样的先用大量无标注中文语料做预训练让模型学会基本的语言规律然后用领域相关的数据做CPT把通用能力往特定方向拉一拉接着用指令数据做SFT让模型学会听懂人话再用PEFT做参数高效微调省显存省时间之后用蒸馏把大模型的能力迁移到小模型上最后用DPO做偏好对齐让输出更符合人类喜好。每一步都有它存在的理由缺了哪一环模型的表现都会差一截。2. 预训练让模型先学会“说人话”2.1 预训练到底在训练什么预训练的本质是让模型通过大量文本学会预测下一个词。这个过程看起来简单但模型在反复预测的过程中会逐渐掌握语法结构、语义关联、常识知识甚至一些简单的推理能力。你可以把它理解成一个小孩在大量阅读中慢慢学会了语言的组织方式虽然他还不会回答问题但他已经知道什么样的句子是“像话的”。Xihe 的预训练目标采用的是自回归语言建模也就是给定前面的词预测下一个词。损失函数就是交叉熵计算预测分布和真实下一个词之间的差异。这个目标函数虽然简单但梯度信号非常密集——每个位置都会产生一个损失值模型参数能得到充分的更新。预训练的数据质量直接决定了模型的基础能力。我用的是一批清洗过的中文通用语料总量大概在几十G的纯文本。清洗流程包括去重、去除乱码、过滤低质量段落、按长度截断等。这里有个经验数据清洗花的时间应该比训练本身还多。我见过太多人随便拿一批数据就开始训结果模型输出全是重复和乱码回头排查发现是数据里混了大量爬虫垃圾。2.2 预训练的关键参数与实操配置预训练阶段有几个参数需要重点关照。首先是学习率我用的峰值学习率是 3e-4配合 warmup 和余弦衰减。warmup 步数设成总步数的 5% 左右让模型在初期不要被太大的梯度冲击。batch size 方面我用的是梯度累积来模拟大 batch实际单卡 batch size 设成 16累积 8 步等效 batch size 128。序列长度设成 512这个长度在中文场景下能覆盖大部分段落。模型结构上Xihe 的层数、隐藏维度、注意力头数需要根据你的算力来定。我用的配置是 12 层、隐藏维度 768、12 个注意力头参数量大概在 1 亿左右。这个规模在单张 24G 显存的卡上跑预训练是可行的但需要开混合精度训练和梯度检查点。如果你显存更小可以把层数降到 6 层隐藏维度降到 512先跑通流程再说。训练过程中要盯紧几个指标训练损失是否稳定下降、验证损失是否过拟合、梯度范数是否爆炸。我一般每 500 步记录一次验证损失如果连续几次验证损失不降反升就要考虑早停或者调小学习率。梯度范数如果超过 1.0说明梯度爆炸了需要加梯度裁剪我设的阈值是 1.0。注意预训练阶段不要过早看生成效果。模型在损失降到 3.0 以下之前生成的文本基本都是乱码这是正常的。我见过有人在损失还在 5.0 的时候就急着看输出然后怀疑自己代码写错了白白浪费好几天。2.3 预训练中的常见坑与排查第一个坑是损失不下降。如果训练损失在最初几百步几乎不动大概率是学习率太小或者数据有问题。我的排查顺序是先检查数据加载是否正确、标签是否对齐再检查学习率是否被 warmup 压得太低最后看模型初始化是否有问题。第二个坑是损失突然变成 NaN。这通常是梯度爆炸或者混合精度训练中的数值溢出导致的。解决办法是加梯度裁剪、降低学习率、检查数据中是否有异常长的序列。我在一次实验中遇到 NaN排查了半天发现是某条数据里混入了一个超长的不间断字符串导致位置编码溢出。第三个坑是过拟合。小模型在大量数据上预训练时如果数据量不够大或者重复采样过多验证损失会先降后升。这时候需要增加数据量、加 dropout、或者减小模型规模。我一般会在预训练阶段加 0.1 的 dropout效果比较稳。3. CPT把通用能力往领域方向拉一把3.1 CPT 的定位与适用场景CPT 是 Continued Pre-Training 的缩写中文叫继续预训练。它和预训练的区别在于预训练是从随机初始化开始CPT 是在已经预训练好的模型基础上用领域相关的数据继续训练。你可以把它理解成一个人已经学会了通用语言现在要让他去读某个专业的文献读多了自然就懂这个领域的术语和表达了。CPT 适合什么场景如果你的目标领域有大量无标注文本比如医疗病历、法律文书、金融研报而这些文本的表达方式和通用语料差异较大那 CPT 就很有必要。我这次做的 CPT 是用一批技术社区的文章和文档目的是让模型更懂技术领域的表达习惯。CPT 的数据量不需要像预训练那么大但质量要求更高。我用的数据大概在几G左右全部经过人工抽检确保没有低质量内容。学习率要比预训练小一个数量级我用的是 3e-5避免把预训练学到的通用能力冲掉。3.2 CPT 的训练策略与效果评估CPT 的训练策略和预训练基本一致也是自回归语言建模但有几个细节需要调整。首先是数据混合比例我一般会把领域数据和通用数据按 7:3 混合防止模型在领域数据上过拟合而丢失通用能力。其次是训练步数CPT 不需要跑太多步通常 1 到 2 个 epoch 就够了跑多了反而会过拟合。效果评估方面我主要看两个指标一是领域数据的验证损失是否下降二是通用数据的验证损失是否保持稳定。如果领域损失降了但通用损失升了说明模型在“偏科”需要调整数据混合比例。我还会用一些领域相关的问题做人工评估看看模型的回答是否更专业了。实操心得CPT 之后一定要做一次通用能力的回归测试。我有一次 CPT 跑得太猛模型在技术问答上表现很好但日常对话变得很生硬后来把领域数据比例降到 5:5 才恢复正常。4. SFT让模型学会听懂指令4.1 SFT 的数据构造与格式设计SFT 是 Supervised Fine-Tuning 的缩写中文叫有监督微调。这个阶段的目标是让模型学会按照指令回答问题。预训练后的模型虽然会说话但它不知道什么时候该回答什么SFT 就是教它这个规矩。SFT 的数据格式通常是“指令-输入-输出”的三元组。我构造的数据包括技术问答、代码解释、文档摘要等类型。每条数据都经过人工审核确保答案准确、表达自然。数据量不需要很大几千到几万条就够关键是质量要高。我这次用了大概 2 万条指令数据覆盖了十几个技术子领域。数据格式上我用的模板是### 指令 {instruction} ### 输入 {input} ### 回答 {output}这个模板在训练时会拼成一个完整的序列损失只计算回答部分指令和输入部分的损失被 mask 掉。这样模型学到的就是“给定指令和输入生成回答”。4.2 SFT 的训练细节与注意事项SFT 的学习率比 CPT 还要小我用的是 1e-5。训练轮数一般 3 到 5 个 epoch太多会过拟合。batch size 可以比预训练小一些我用的是 32。序列长度根据数据长度分布来定我设的是 1024覆盖大部分样本。训练过程中要关注回答部分的损失而不是整体损失。因为指令和输入部分的损失被 mask 了整体损失会偏低参考意义不大。我一般看回答部分的困惑度降到 2.0 以下基本就收敛了。SFT 阶段最容易出现的问题是灾难性遗忘。模型在学会指令跟随之后可能会丢失一些预训练阶段学到的知识。缓解办法是在 SFT 数据中混入少量预训练数据或者用较小的学习率、较少的训练轮数。我一般会在 SFT 数据中混 10% 的通用语料效果比较稳。注意SFT 数据的质量比数量重要得多。我试过用 10 万条低质量数据训练效果远不如 2 万条精标数据。低质量数据里的错误答案会被模型学进去而且很难纠正。5. PEFT用最少的参数做最高效的微调5.1 PEFT 的核心思路与主流方法PEFT 是 Parameter-Efficient Fine-Tuning 的缩写中文叫参数高效微调。它的核心思路是不更新模型的所有参数只更新一小部分额外添加的参数或者只更新部分原有参数。这样做的目的是省显存、省时间、减少过拟合风险。主流的 PEFT 方法有几种。LoRA是在原有权重旁边加一个低秩矩阵训练时只更新这个低秩矩阵。Prefix Tuning是在输入前面加一段可学习的向量。Adapter是在每层后面插入一个小型网络。我这次用的是 LoRA因为它实现简单、效果稳定、推理时可以直接合并回原模型。LoRA 的核心参数是秩和缩放系数。秩决定了低秩矩阵的维度秩越大表达能力越强但参数量也越大。我用的秩是 8缩放系数是 16。这个配置在 1 亿参数的模型上可训练参数只有几百万显存占用降低了 60% 以上。5.2 LoRA 的实操配置与效果对比LoRA 的配置需要注意几个点。首先是应用位置我一般把 LoRA 加在注意力层的 query 和 value 投影上这两个位置对下游任务的影响最大。其次是学习率LoRA 的学习率可以比全量微调大一些我用的是 3e-4。最后是训练轮数LoRA 收敛比较快一般 2 到 3 个 epoch 就够了。效果对比方面我在同样的 SFT 数据上跑了全量微调和 LoRA 微调。全量微调的显存占用是 22GLoRA 是 8G。训练时间上全量微调一个 epoch 要 2 小时LoRA 只要 40 分钟。效果上LoRA 在大部分任务上能达到全量微调 95% 以上的水平个别任务甚至更好因为 LoRA 的正则化效果减少了过拟合。实操心得LoRA 的秩不是越大越好。我试过秩 64效果和秩 8 差不多但参数量多了 8 倍。对于小模型来说秩 4 到 16 之间就够用了。另外LoRA 的缩放系数和秩要配合调整一般缩放系数是秩的 2 倍左右比较合适。6. 蒸馏把大模型的能力迁移到小模型6.1 蒸馏的基本原理与实现方式蒸馏的核心思想是让一个小模型学生模型去模仿一个大模型教师模型的输出分布。教师模型的输出不仅包含正确答案的信息还包含错误答案之间的相对概率关系这些信息被称为“暗知识”。学生模型通过学习这些暗知识能在参数量小得多的情况下达到接近教师模型的效果。蒸馏的实现方式有两种。软标签蒸馏是让学生模型去拟合教师模型的 softmax 输出分布通常用温度参数来平滑分布。硬标签蒸馏是让学生模型直接拟合教师模型的最终输出。我这次用的是软标签蒸馏温度设成 2.0损失函数是 KL 散度。教师模型我用的是一个更大的预训练模型学生模型就是前面训练好的 Xihe。蒸馏数据用的是 SFT 数据加上一批无标注文本。训练时教师模型和学生模型同时前向传播学生模型的损失是蒸馏损失和真实标签损失的加权和权重我设的是 0.7 和 0.3。6.2 蒸馏过程中的关键参数与效果评估蒸馏的关键参数是温度和损失权重。温度越高教师模型的输出分布越平滑暗知识越丰富但太高会导致信息模糊。我试过温度 1.0、2.0、5.0发现 2.0 效果最好。损失权重方面蒸馏损失占主导真实标签损失起辅助作用我用的 0.7:0.3 比较平衡。效果评估上我对比了蒸馏前后学生模型在测试集上的表现。蒸馏前学生模型的准确率是 78%蒸馏后提升到 84%接近教师模型的 87%。参数量上学生模型只有教师模型的十分之一推理速度快了 8 倍。这个 trade-off 在实际部署中非常划算。注意蒸馏的效果很大程度上取决于教师模型的质量。如果教师模型本身就不够好学生模型学到的也是半吊子。另外蒸馏数据的分布要和目标任务匹配否则学生模型会在无关数据上浪费容量。7. DPO让模型输出更符合人类偏好7.1 DPO 与 RLHF 的关系及优势DPO 是 Direct Preference Optimization 的缩写中文叫直接偏好优化。它的目标是让模型的输出更符合人类偏好比如更有帮助、更安全、更礼貌。传统的 RLHF 需要先训练一个奖励模型再用强化学习优化策略模型流程复杂且不稳定。DPO 直接用一个偏好损失函数来优化模型省掉了奖励模型和强化学习环节实现简单、训练稳定。DPO 的数据格式是“指令-优选回答-劣选回答”的三元组。我构造的数据包括技术问答中的好答案和差答案对比比如一个回答详细准确另一个回答含糊错误。数据量不需要很大几千条就够关键是偏好标注要一致。DPO 的损失函数核心是让模型对优选回答的概率相对劣选回答的概率更高。具体来说它计算模型对优选回答和劣选回答的 log 概率之差然后通过 sigmoid 函数转换成损失。这个损失函数的设计使得模型不需要显式的奖励信号就能从偏好数据中学习。7.2 DPO 的训练配置与效果观察DPO 的学习率要比 SFT 小我用的是 5e-6。训练轮数一般 1 到 2 个 epoch太多会导致模型过度拟合偏好数据输出变得极端。batch size 我用的是 16因为 DPO 需要同时计算优选和劣选回答的概率显存占用比 SFT 大。训练过程中要关注偏好准确率也就是模型对优选回答的评分高于劣选回答的比例。这个指标在训练初期会快速上升然后趋于稳定。我一般训练到偏好准确率 85% 以上就停再训下去提升有限反而可能损害模型的通用能力。DPO 之后模型的输出会变得更“讨喜”但也要注意不要过度优化。我遇到过 DPO 训练太久导致模型变得过于保守对任何问题都给出模棱两可的回答。后来把训练轮数从 3 降到 1问题就解决了。实操心得DPO 的偏好数据质量比数量重要。我试过用自动构造的偏好数据训练效果远不如人工标注的。人工标注的偏好数据虽然少但一致性高模型学到的偏好更准确。另外DPO 之后最好再做一次通用能力的回归测试确保模型没有在偏好数据上过拟合。8. 全流程串联与工程化建议8.1 各阶段的数据与模型流转把整个流程串起来看数据流转是这样的预训练用大规模无标注语料产出基础模型CPT 用领域无标注语料产出领域基础模型SFT 用指令数据产出指令跟随模型PEFT 在 SFT 模型上做参数高效微调产出轻量适配模型蒸馏用教师模型和学生模型产出压缩后的小模型DPO 用偏好数据产出对齐后的最终模型。每个阶段的产出都可以单独使用也可以作为下一阶段的起点。我一般会把每个阶段的模型都保存下来方便对比和回滚。模型保存时要注意保存优化器状态和学习率调度器状态方便断点续训。工程化方面我建议用配置文件管理每个阶段的超参数用实验跟踪工具记录训练曲线和评估指标。我用的是一套自己写的轻量级训练框架支持混合精度、梯度累积、梯度检查点、分布式训练等功能。如果你不想自己写也可以用现成的训练框架但要注意框架的版本兼容性。8.2 资源规划与时间估算资源规划方面预训练是最耗资源的阶段。我用的单卡 24G 显存1 亿参数模型序列长度 512batch size 16梯度累积 8训练 10 万步大概需要 3 天。CPT 和 SFT 各需要半天到一天。PEFT 和 DPO 各需要几个小时。蒸馏取决于教师模型的大小我用的是 10 亿参数的教师模型蒸馏一遍大概需要一天。如果你资源有限可以按这个优先级来先跑 SFT 和 PEFT这两个阶段对资源要求最低效果也最直观。然后跑 DPO提升输出质量。最后再考虑预训练和 CPT这两个阶段需要大量数据和算力。蒸馏可以在有教师模型的情况下做没有的话可以跳过。注意不要一上来就想着跑全流程。我建议先把 SFT 和 PEFT 跑通确认整个训练和评估流程没问题再逐步往前扩展。我见过太多人一上来就搞预训练结果数据没准备好、代码有 bug、算力不够折腾几周都没跑出一个能用的模型。9. 常见问题速查与避坑指南9.1 训练过程中的典型问题问题现象可能原因排查方法解决方案损失不下降学习率太小、数据有问题、模型初始化异常检查数据加载、学习率配置、初始化代码调大学习率、清洗数据、重新初始化损失变 NaN梯度爆炸、混合精度溢出、异常长序列检查梯度范数、数据长度分布加梯度裁剪、降低学习率、过滤异常数据验证损失上升过拟合、数据重复、模型太大检查训练轮数、数据去重情况早停、增加数据、减小模型生成重复文本解码策略问题、训练不充分检查解码参数、训练损失调温度、重复惩罚、继续训练灾难性遗忘学习率太大、训练轮数太多对比微调前后的通用能力降低学习率、减少轮数、混入通用数据9.2 独家避坑技巧第一个技巧是数据去重一定要彻底。我在一次预训练中发现模型对某些句子记得特别牢后来排查发现是数据里有大量重复段落。去重不能只用精确匹配还要用模糊匹配和语义去重。我用的是一套基于 MinHash 的去重流程效果不错。第二个技巧是学习率要跟着 batch size 调。如果你增大了 batch size学习率也要相应增大否则训练会变慢。经验法则是学习率与 batch size 的平方根成正比。我一般会先在小 batch 上试出合适的学习率再按这个比例放大。第三个技巧是评估要分阶段做。不要等全流程跑完才评估每个阶段结束后都要做一次评估确认这个阶段的目标达成了。比如 SFT 之后要测指令跟随能力DPO 之后要测偏好准确率。这样一旦某个阶段出问题能快速定位。第四个技巧是保存检查点要勤。训练过程中随时可能因为各种原因中断保存检查点能让你从最近的状态恢复。我一般每 1000 步保存一次同时保留最近三个检查点。检查点要包含模型权重、优化器状态、学习率调度器状态和训练步数。第五个技巧是推理测试要用真实场景。训练时的评估指标再好也不代表实际使用效果好。我一般会在每个阶段结束后用一批真实用户问题做人工评估看看模型的回答是否自然、准确、有帮助。这个环节能发现很多自动指标发现不了的问题。10. 我个人的一些体会这套流程我前前后后跑了不下十遍每次都有新的收获。最大的体会是语言模型的训练是一个系统工程每个环节都会影响最终效果。预训练的数据质量决定了模型的基础能力上限SFT 的数据质量决定了模型的指令跟随能力DPO 的数据质量决定了模型的输出风格。任何一个环节偷懒最终效果都会打折扣。另一个体会是小模型有小模型的优势。虽然小模型的知识容量有限但它在特定任务上可以做得很好而且推理速度快、部署成本低。我训练的这个 Xihe 小模型在技术问答任务上已经能满足大部分日常需求响应速度比大模型快得多。最后想说的是不要怕踩坑坑踩多了就成路了。我一开始跑预训练的时候损失不下降、梯度爆炸、显存溢出什么问题都遇到过。但每次解决问题之后对模型训练的理解就深了一层。现在回头看那些踩过的坑才是最宝贵的经验。
上一篇/下一篇内容由系统自动关联
返回资讯列表 →