ARTICLE DETAIL

资讯详情

深耕网站建设、视觉设计与SEO优化的一线实战洞察。

从CPT到DPO:单卡训练小语言模型Xihe的完整实践

从CPT到DPO:单卡训练小语言模型Xihe的完整实践 1. 项目概述与整体设计思路Xihe 是我今年一直在推进的一个小语言模型项目。从选底座、整理语料、做持续预训练CPT开始再到指令微调SFT、参数高效微调PEFT、模型蒸馏最后用 DPO 做偏好对齐整条链路都走了一遍。这篇文章就是这段实操记录的完整版里面包括我调整过的参数、踩过的坑以及为什么在每一步选那种方案希望能给打算在有限算力下拥有私有小模型的人一点参考。这里说的“从零训练”并不是指随机初始化一个模型然后硬训几万亿 token而是指不直接依赖现成的大模型 API自己搭起一条“数据准备 → 继续预训练 → 指令微调 → 蒸馏 → 偏好对齐”的完整流水线。我管这叫“从零”因为每一阶段的数据、脚本和评估方式都得自己设计完全没有捷径可走。1.1 核心需求解析我的目标其实很明确在有限的单卡环境下做出一个参数量不大但能真正用于垂直领域问答的模型。这听起来简单实际牵涉到的东西却不少。首先得确定底座。既然要跑在消费级 GPU 上参数量就不能太大。我选了约 1B 参数的 decoder-only 架构不是因为它最强而是因为后续任何实验——CPT、SFT、DPO——都能在合理时间内跑完。如果一上来就用 7B 甚至 13B单卡训练会非常痛苦而且迭代速度直接拖垮整个项目。其次是数据问题。一个小语言模型最吃亏的地方就是知识容量和泛化能力都有限所以语料的质量比数量更重要。我花了接近两周时间处理领域语料包括去重、过滤、配比和清洗实际训练消耗的时间都比这个短。很多朋友以为训练才是最核心的但真正跑下来才发现语料工程才是决定模型上限的那只手。最后是全流程解耦。我没有把 CPT、SFT、蒸馏、DPO 揉成一个大的训练任务而是拆成四个相对独立的阶段每个阶段只改动一部分参数或数据。这样出了问题能精准定位而且每一步产出的 checkpoint 都可以保留方便回滚对比。1.2 为什么选择 CPT 而不是从零预训练从零预训练是一个很浪漫的想法但现实非常骨感。要从随机权重开始训练一个 1B 模型至少需要上千亿 token 才能做到基础的通顺和知识覆盖这个量级的数据获取和算力消耗个人和小团队基本扛不住。所以我做了个折中找一个质量还不错的通用中文预训练模型作为起点再用自己的领域数据做持续预训练。这样做有三层好处。第一通用语言能力和基本知识都已经具备我只需要让模型适应目标领域的数据分布和词汇表达训练成本大幅下降。第二底座模型大概率有比较成熟的 tokenizer 和稳定的训练设置我能踩着前人调好的参数前进而不是从头摸索学习率、warmup 这些细节。第三CPT 阶段可以用很小的学习率我最后定在 1e-4 量级做微调式更新不容易把原有能力冲掉后续 SFT 才有足够的基础可以依赖。很多做过 NLP 的朋友可能会问直接用一个通用中文底座去做 SFT 行不行答案是可以但领域效果会差一截。特别是一些专业术语、行话、独特的表达方式通用模型见到的少了SFT 阶段就算见过也很难真正内化。CPT 就是先把这些领域语言结构“预装”进去让后续微调事半功倍。1.3 方案选型背后的取舍逻辑我在架构和训练框架上做了几次取舍简单总结一下逻辑。底座方面我对比过几类中文预训练模型包括 RoBERTa 风格、GPT-2 风格和近年更新的开源语言模型。RoBERTa 风格长于理解任务但生成能力偏弱GPT-2 类虽然老但结构干净迁移到新框架容易。最终我选了一个结构接近 GPT-2 的中文预训练模型作为起点因为后续要做生成型对话和指令微调自回归式的生成能力更重要。训练框架上我用了 Hugging Face 的 transformers 和 PEFT 库再配合 deepspeed 做显存优化。PEFT 是关键的转折点前期用全参数微调做了基准实验发现显存压力和训练时间都难以接受后来切到 LoRA显存占用直接掉了一大截效果却几乎没有差别。这也是为什么我把 PEFT 单独拿出来讲——后面 SFT 和 DPO 两阶段我都依赖它。蒸馏则放在 PEFT 之后。我一开始想用蒸馏替代 SFT认为把大模型知识灌进小模型就行后来发现单纯蒸馏缺少指令遵循能力模型生成的内容虽然通顺但不太听话。最终把蒸馏设计成“先 SFT 学会服从再蒸馏压缩知识最后 DPO 对齐偏好”的流程效果才稳定下来。2. 持续预训练CPT实操持续预训练是整个项目的第一场硬仗。这一阶段的目标不是让模型学会某个具体任务而是让它吸收领域数据中的知识、术语和表达方式为后面的指令微调打好底座。2.1 领域语料筛选与清洗CPT 阶段我最重视的就是数据甚至可以说数据质量直接决定了模型是“聪明”还是“死板”。我从几个不同的公开来源收集了领域语料包括论文摘要、技术博客、产品文档、问答社区帖子等。但原始数据基本都不能直接用必须经过几轮清洗。第一轮是格式清洗。去掉 HTML 标签、Markdown 残留符号、乱码字符、广告和重复的版权声明等。这部分我写了个 Python 脚本用正则表达式和简单的规则批量处理。注意不要过度清洗比如代码块里的缩进、表格分隔符如果一刀切删掉反而会影响模型对结构文本的理解。第二轮是质量过滤。我按几个规则打分文本长度是否过短、句子是否完整、是否包含大量非中文或数字噪声、是否有明显的重复模式。得分太低的直接扔掉。特别提一下重复问题公开语料里存在大量重复段落如果不做去重模型会把某些高频表达固化为唯一输出生成时非常呆板。第三轮是配比设计。我不会把领域语料单独喂进去而是按约 631 的比例混合领域数据、通用中文语料和少量代码文本。这样处理是为了防止“灾难性遗忘”——如果模型一直在看领域文档很快就会忘记通用对话的表达方式后续 DPO 阶段会非常难拉回来。清洗完成后我大约保留了 12GB 的纯文本数据。这个量对从零预训练来说很小但对 CPT 来说已经足够关键在于质量够高、覆盖均匀。2.2 训练目标与关键参数设置CPT 阶段的训练目标和基座模型保持一致我用的基座是自回归架构所以训练目标就是下一个 token 预测。这里有个容易出错的地方很多人以为 CPT 应该换成 MLM掩码语言模型任务但那样会破坏模型的生成连贯性。除非你的底座本身就是 RoBERTa 这类 encoder-only 模型否则不要随便更换训练任务。训练参数上我参考了社区里对小模型做 CPT 的常见设置再结合自己卡上的显存做调整参数取值说明学习率1e-4带余弦衰减比全量预训练低避免冲掉原有知识warmup 步数1000 步稳定早期训练batch size512 条约 0.5M tokens通过梯度累积实现训练轮数0.5 epoch意思是对领域语料只过一半混合精度bfloat16显存和速度兼顾最大序列长度2048适合文档类语料有人会疑惑为什么训练轮数只有 0.5 epoch这不是让模型没学完吗其实持续预训练和普通预训练不一样领域数据量不大时重复多轮非常容易过拟合模型会逐渐丢掉通用能力。我只让模型看到部分领域数据配合低学习率做“浅层吸收”后续 SFT 还能进一步强化关键模式。训练过程中我实时观测的是 loss 曲线但说实话loss 并不是唯一标准。这个阶段我更关心领域样本上的困惑度变化以及模型在几个固定 prompt 下生成文本的流畅度。如果 loss 在降但生成结果出现大量重复那说明学习率可能太高或数据去重不彻底要赶紧停下来调整。2.3 单卡训练的资源估算与调优我用的是 24GB 显存的消费级显卡1B 参数模型在全参数训练时是放不下的所以 CPT 阶段我就开始用 LoRA。可能有朋友会问CPT 阶段就用 PEFT那不叫真正意义的继续预训练了吧我的看法是对 1B 这样的小模型来说LoRA 微调已经能对领域知识做相当不错的适配尤其在数据量不大、迭代次数很多的场景下效果接近全量微调但显存和速度的收益非常明确。如果非要做全参数 CPT我建议把模型拆成冻结和解冻两部分或者直接上 deepspeed stage-2。我记得在没做任何优化的情况下1B 模型全参数训练光优化器状态就要吃掉好几个 GB再加上激活值和梯度24GB 卡基本被压到极限batch size 只能设成 1。切到 LoRA 之后模型参数冻结只有低秩矩阵参与更新显存压力一下降了很多batch size 能提到 8 甚至 16配合梯度累积训练速度反而更快。有人问怎么判断训练多久合适。我一般看两个信号一是领域数据的 loss 降到通用数据 loss 的 70%-80% 左右二是拿模型跑三五个领域相关 prompt看输出有没有出现原文里的专业术语和逻辑。只要这两点都达标我就停掉 CPT不再为了让 loss 降得更低而多烧机器。3. SFT 与 PEFT 指令微调实操CPT 做完之后模型就像一个“读过很多领域资料但不会回答问题”的人知识在肚子里一问就懵。SFT 就是教它如何把知识组织成符合用户期待的回复。3.1 指令数据集的构建与格式设计我在 SFT 阶段用了大约 8 万条指令数据其中大部分来自人工标注再配合一部分从社区收集的高质量对话最后用模型辅助生成了一些扩写样本。数据量看起来不大但这个数字已经是精筛后的结果宁缺毋滥。指令数据的格式直接影响后面训练和推理的效果。我先定义了一种统一模板|im_start|user 我是 XX 场景下的运营人员请帮我总结这段内容的要点 {input_text}|im_end| |im_start|assistant {response_text}|im_end|这种带特殊分隔符的格式比简单的 “Question: ... Answer: ...” 要清晰得多模型在推理阶段只要看到|im_start|user就知道接下来是用户输入看到|im_start|assistant就开始生成回复不会混淆角色。构造数据时我特别强调两个原则。第一是输入多样性同一个意图尽量用不同句式表达比如“帮我总结”“请提炼要点”“简单概括一下”这样模型不会把某个句式当作唯一触发条件。第二是答案规范性每条 answer 都得经过人工审查不出现含糊其辞、前后矛盾或安全风险内容。很多朋友为了凑数据量把模型生成的答案直接丢进训练集结果越训质量越差这个坑我踩过后面在问题清单里详细说。3.2 LoRA 参数选择与训练过程SFT 阶段我用的是 LoRA。之前用全参数微调跑了 2 个 epoch效果不错但显存峰值很高而且每次调整数据都要重新训练成本太大。LoRA 的核心思路很直观冻结原始权重只训练注入到模型中的低秩矩阵在推理时又能把增量合并回原权重。我把它理解为给模型做“外挂微调”训练时只需要更新很小一部分参数。具体配置如下参数取值目标模块q_proj, k_proj, v_proj, o_projrank16alpha32dropout0.05学习率3e-4batch size32经过梯度累积训练轮数5 轮但设置了早停rank 是 LoRA 最重要的超参数之一。我试过 8、16、32 三档rank8 时模型学得稍慢rank32 时训练时间明显增加且没有看到效果提升最终定在 16。alpha 一般取 rank 的 2 倍也可以直接调但我不建议一上来就把 alpha 拉得太大否则可能引入数值不稳定性。训练过程中我监控两个指标训练集 loss 和验证集 loss。SFT 非常容易过拟合尤其在小数据集上经常训练集 loss 一路下降验证集 loss 却在某个点开始反弹。我记录了每个 epoch 的 checkpoint按验证集 loss 最低的那个 epoch 作为最终模型而不是最后一个 epoch。3.3 SFT 效果评估与迭代策略模型训完我从来不看单一指标就拍板。常规的做法是准备一套固定的评测集包含 30 个典型问题覆盖生成、提取、总结、纠错等真实场景每次调完数据就跑一遍。评测分两头看。自动评估方面我用 ROUGE 和 BLEU 这类指标做参考但它们很难反映生成内容的实际质量只能看有没有跑偏。真正让我信服的还是人工打分每个问题让两位同学独立打分维度包括“信息正确性”“表达流畅度”“指令遵循度”最后取平均。用这套办法我能清楚看到某批数据改动到底带来的是净提升还是错觉。如果评测分数不够理想我会回头检查数据而不是立刻加训练轮数。最常见的病根是数据配比失衡比如“总结类”样本太多模型就会倾向于把所有问题都答成总结格式或者 answer 里出现大量模板化开头“根据您的问题我的回答如下”模型也会学成复读机。这个阶段的核心是数据迭代训练只是放大器。4. 模型蒸馏实操蒸馏的目的很朴素让一个小模型去模仿一个大模型的行为。我之所以做蒸馏是因为最后要部署的是一个更小的模型大约 300M 参数比 1B 的底座小不少。如果直接拿 SFT 后的 1B 模型去部署推理速度不够快显存占用也太高根本不适合做实时的在线服务。4.1 蒸馏方案的定位为什么放在 SFT 之后很多人把蒸馏理解成“用大模型生成数据来训练小模型”这只说对了一半。蒸馏更内核的东西是让学生模型学习教师模型的概率分布而不仅仅是学习它的输出文本。我在设计流程时把蒸馏放在 SFT 之后而不是之前原因是如果先蒸馏再 SFT学生模型学到的是一个大模型在通用指令数据上的行为虽然知识密度高但对具体任务指令的理解还很弱。反过来先做 SFT 再蒸馏教师模型已经是一个会“按指令做事”的模型学生模仿的就是它的完整行为模式包括格式、语气和任务切换能力。实验对比下来后者的效果更加稳定。蒸馏阶段用到的教师模型是一个 7B 级别的开源模型跑在另一台设备上。学生模型用的是 300M 的小底座重新初始化。有些方案会直接拿 SFT 后的 1B 模型当学生再往 300M 压缩但我试过之后觉得跨度过大效果不稳。拆成两步——1B 先蒸馏到 0.5B再蒸馏到 0.3B——虽然麻烦但每一步都得到了更好的结果。4.2 软标签、温度与蒸馏损失设计标准的知识蒸馏损失包含两部分。第一部分是让学生模型对教师模型的软标签输出建模用 KL 散度衡量两个概率分布的差异。第二部分是让学生模型对真实标签做常规的交叉熵损失防止学生模型完全被教师模型的错误带偏。这两部分的权重我最初设定为 73后来发现偏重软标签时学生模型学到的“风格”更多偏重真实标签时“事实正确性”更稳最后调成 64。温度 T 是关键参数。教师模型输出的 logits 经过一个带温度 T 的 softmax 后变成更平滑的概率分布小的 T 会让分布尖锐接近硬标签大的 T 会让分布扁平突出相似类别之间的相对差异。我试了 1.0 到 8.0 几档最后发现 4.0 效果不错太低时学生学到的东西太“窄”太高时分布过于平均反而模糊了重要信息。学生模型的损失可以写成L alpha * KL(softmax(teacher_logits / T) || softmax(student_logits / T)) (1 - alpha) * CE(student_logits, ground_truth)这个公式不算复杂但有一点要注意KL 散度中的温度 T 不会自己消失训练时需要对教师和学生 logits 都除以 T而且最终推理时不能再除以 T。如果忘了把温度重标定回去学生模型的输出会变得非常平滑生成各种含糊不清的文本。4.3 蒸馏数据扩充与训练稳定性蒸馏训练的数据主要来自两个渠道。一是已有的 SFT 指令数据直接让学生模型在相同问题上模仿教师模型的回答。二是教师模型在新 prompt 上的生成数据我会刻意构造一些真实用户可能问但原数据里没有覆盖的问题让教师模型回答之后加入训练集。这里有个实操技巧让教师模型生成答案时temperature 要适当调低一点我一般设置在 0.7 左右避免采样出太离题的文本。但同时还需要做一次质量过滤——如果教师模型对某个 prompt 的输出明显混乱或安全合规上不放心就直接丢掉不要进训练集。用质量不高的数据做蒸馏等于把错误知识放大教给学生。蒸馏训练本身还算稳定主要问题是显存。因为学生模型和教师模型要同时在前向传播中计算两个模型都会占显存。我的处理方式是先把教师模型的 logits 离线算出存成文件训练时直接读取不再跑教师模型。这样一来训练阶段只需加载学生模型和预计算的标签显存占用低了一大截训练速度也快了很多。5. DPO 偏好对齐实战模型到了这个阶段已经能做到“知识在脑、指令顺手、身形轻巧”了但还有一个隐性问题模型可能会生成安全上不合规、或者憋着不输出用户真正想要的内容。DPO 就是对模型进行偏好对齐的实用手段让模型学会什么是更好的回答。5.1 DPO 的原理与和 RLHF 的对比DPO 的全称是 Direct Preference Optimization直接偏好优化。它的核心思路是不需要为模型训练一个奖励模型也不需要在线做强化学习采样而是直接把偏好数据转换成损失函数脱离对 RLHF 复杂链路和超多超参数的依赖。我以 RLHF 做对比来理解 DPORLHF 要训练一个 reward model 来给回答打分再通过 PPO 让模型学着最大化分数过程繁琐且对算力要求很高DPO 则把“模型更偏好哪个回答”这一偏好对直接作为监督信号通过一个解析解计算出最优策略的更新方向。听起来很玄乎实际操作中它就是样本对的形式每个样本都包含“可接受回答”和“不可接受回答”两种版本。当然 DPO 也不是完全没有代价。它最大的前提是训练数据必须优质且偏好方向明确。如果两个回答在质量上差不多或者其中一个只是风格不同最终模型很容易产生波动。所以偏好数据的构建我格外谨慎后面专门写一节。从工程角度看DPO 默认情况下是对整个模型权重做更新的如果你的模型已经经过前面几轮微调直接全参数 DPO 会有灾难性遗忘风险因此我把 DPO 也放在 PEFT 框架之下——仍然用 LoRA。这样既让模型学习偏好信号又保证主体参数几乎不动。5.2 偏好数据集的构建思路我先花了大量精力构建偏好对。每条偏好对格式包括三部分一个 prompt、一个更优的回答chosen、一个更差的回答rejected。那么这些“更差回答”从哪来一部分来自之前 SFT 模型在不同参数下产生的输出一部分来自不同温度采样导致的不理想结果还有一部分是人工标注结果的对比。值得说的是我构建偏好对时参考了模型自身的判断如果一个回答出现事实错误、答非所问、语气别扭、隐含拒绝用户请求之类的问题就会被标记为 rejected。chosen 回答通常是经过人工核实、事实正确、表达清晰且安全的版本。数据里我还故意保留了一些难度比较高的例子。比如 prompt 本身模棱两可时chosen 回答会主动向用户确认需求而不是蠢答一通rejected 回答则写成猜测式、含糊式。这样模型能学到的不只是“说什么好”还有“在信息不足时该怎么应对”。偏好对的数量我用到了大概 5 万条这个量级在 DPO 训练里属于很小的但因为质量很高实际效果非常好。多余的低质数据不仅没用还会引入噪声。5.3 DPO 训练参数与常见问题DPO 训练有一个非常经典的超参数 beta它控制对参考模型的依赖程度。beta 越大模型越不愿意偏离 SFT 阶段的参考模型beta 越小模型越积极地适应偏好数据。社区常见做法是 beta 取 0.1 到 0.5我在这个项目里最终用的是 0.3。太大会让偏好学习很微弱太小则容易直接把模型训“崩”——输出开始变漂浮甚至回答越来越短明显失去生成多样性。学习率也要比 SFT 阶段低不少。我用了 1e-6 做 LoRA 微调训练轮数控制在 1 到 2 轮之内。DPO 的论文和社区经验都说得很清楚DPO 过度训练会让模型产生“奖励退化”现象就是偏好数据上的得分一直涨但通用能力大幅下降。务必保留 checkpoint在每个 epoch 结束之后用真实评测集做一次人工抽测。训练过程中我常碰到的三个问题这里直接给速查表症状可能原因处理方式训练 loss 快速降低评测却变差偏好对本身质量差或分布太窄重新筛选偏好对增广 prompt 多样性模型回答变短、让步多、无主见beta 过大或偏好数据中 rejected 太多下调 beta调高 chosen 回答的信息量训练后模型出现重复句式学习率偏高、数据量过大降低学习率提前早停在实际操作里我发现一个特别容易被忽略的点DPO 阶段一定要用参考模型计算原始 log prob。这个参考模型不是教师模型而是 SFT 完成后的那个模型。你需要把参考模型的参数冻结在训练时和当前模型一起计算对数概率比值。如果忘了冻结DPO 就变成一个普通的对比学习效果会打折扣。6. 常见问题与避坑清单整个项目跑下来我遇到过的坑加起来可以开一个吐槽帖了。这一章专门把最有代表性的问题和排查思路写给后来者希望能帮你少走弯路。6.1 损失函数不降或 NaN 的排查训练刚开始时最容易出诡异问题。有一天 CPT 阶段 loss 在 50 步内暴涨到一个不可思议的值紧接着就变 NaN 了后来排查了半天发现是数据里有极长行计算注意力时 logits 溢出。这是一个非常典型的数值问题输入长度参差不齐时如果采用了错误的 padding 或位置编码策略模型在最后几步就能把数值推向极端。排查思路一般按顺序来第一步看学习率是不是太大尤其对 LoRA 这类低秩结构学习率需要比全参数微调更谨慎第二步看混合精度有没有溢出的风险bfloat16 比 float16 更稳第三步检查数据中是否有异常片段比如连续几千个数字字符、全角半角混乱、非法 token 组合。数据清洗才是根治但紧急情况下把 max sequence length 调小也能快速把模型从 NaN 边缘拉回来。建议在训练脚本里加上 loss 值监控和 NaN 自动停下save 前一个 checkpoint 的副本。这个机制救了我几次否则一个晚上全白跑。6.2 过拟合与评估失真领域数据通常不多所以所有阶段都容易过拟合SFT 尤为严重。模型在训练集上跑得很好一到新数据上就显出原型复述原句、输出模板话术、丢失领域细节。我的判断标准非常简单——拿一个训练中从未见过的提问去问它如果答案还带着训练集里的原句片段那基本就是背下来了得赶紧降低训练轮数。评估指标失真也是一个密切相关的问题。ROUGE 这类指标在“总结类”任务上看起来很高不代表生成质量好。我原以为模型已经能打 80 分结果人工评测只给了 55 分原因是模型经常优缺漏或者强行拼接。从那以后我再也不敢单独依赖自动指标每个阶段都至少跑一组人工评测。另外要注意评测集不要和训练集重叠。很多情况下看起来公平的评测里面其实混着训练集样本模型分数虚高真实上线效果惨不忍睹。建立评测集时务必做去重。6.3 蒸馏与 DPO 阶段的效果倒退蒸馏和 DPO 都是“越优化越可能倒退”的阶段。模型蒸馏时学生模型可能完美模仿了教师模型的表达风格但失去了事实稳定性DPO 时模型可能学会了“讨好”偏好数据但输出范围明显收窄。我的解决办法是保留每个阶段的基座SFT 版、蒸馏版、DPO 版全部存下来。评测时把三版模型同时比对而不是只看最新版。如果 DPO 版本在通用评测集上明显下跌但偏好任务得分上升我会衡量产品需求后决定是否回退。这里有个小技巧无论是蒸馏还是 DPO最终模型都可以和上一阶段的模型做“模型融合”或参数平均比如把 DPO 模型和 SFT 模型的参数按 0.70.3 加权平均。这个方法虽然是土办法但在小模型场景下经常能把倒退拉回来一些代价只是多几次实验和一点点推理代码。6.4 部署与推理阶段的小细节最后提一下部署。300M 的模型虽然小但推理框架的选择依然会影响实际效果。我直接使用 ONNX Runtime 导出模型顺手做 INT8 量化显存占用又降了一截。量化的损失通常不大但要注意 tokenizer 部分也必须对齐很多坑都出在导出后 tokenizer 和模型不一致上。如果你要把模型嵌入到现有服务里我强烈建议做一个“兜底策略”当模型输出的置信度很低时不要硬答可以返回“需要更多信息”或触发旧规则逻辑。小模型的自信心和能力并不总成正比兜底策略往往能提高整体用户体验而不是把所有压力都压在模型上。以我个人经验来说训练一个小语言模型最核心的不是某一步有多高大上而是每一步的输入质量是否配得上训练成本。Xihe 这个项目至今还在迭代中后续我还会尝试更小更快的 checkpoint、更复杂的偏好数据构造以及把整个流水线自动化。如果你也在做类似的实验建议先从一到两个阶段跑通再逐步往上加别一上来就想着 Ablations 全做那样数据、算力和时间都会被吃得很紧。
返回列表