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 脚本,用正则表达式和简单的规则批量处理。注意不要过度清洗,比如代码块里的缩进、表格分隔符,如果一刀切删掉,反而会影响模型对结构文本的理解。
第二轮是质量过滤。我按几个规则打分:文本长度是否过短、句子是否完整、是否包含大量非中文或数字噪声、是否有明显的重复模式。得分太低的直接扔掉。特别提一下重复问题:公开语料里存在大量重复段落,如果不做去重,模型会把某些高频表达固化为唯一输出,生成时非常呆板。
第三轮是配比设计。我不会把领域语料单独喂进去,而是按约 6:3:1 的比例混合领域数据、通用中文语料和少量代码文本。这样处理是为了防止“灾难性遗忘”——如果模型一直在看领域文档,很快就会忘记通用对话的表达方式,后续 DPO 阶段会非常难拉回来。
清洗完成后,我大约保留了 12GB 的纯文本数据。这个量对从零预训练来说很小,但对 CPT 来说已经足够,关键在于质量够高、覆盖均匀。
2.2 训练目标与关键参数设置
CPT 阶段的训练目标和基座模型保持一致,我用的基座是自回归架构,所以训练目标就是下一个 token 预测。这里有个容易出错的地方:很多人以为 CPT 应该换成 MLM(掩码语言模型)任务,但那样会破坏模型的生成连贯性。除非你的底座本身就是 RoBERTa 这类 encoder-only 模型,否则不要随便更换训练任务。
训练参数上,我参考了社区里对小模型做 CPT 的常见设置,再结合自己卡上的显存做调整:
| 参数 | 取值 | 说明 |
|---|---|---|
| 学习率 | 1e-4,带余弦衰减 | 比全量预训练低,避免冲掉原有知识 |
| warmup 步数 | 1000 步 | 稳定早期训练 |
| batch size | 512 条,约 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_proj |
| rank | 16 |
| alpha | 32 |
| dropout | 0.05 |
| 学习率 | 3e-4 |
| batch size | 32(经过梯度累积) |
| 训练轮数 | 5 轮,但设置了早停 |
rank 是 LoRA 最重要的超参数之一。我试过 8、16、32 三档,rank=8 时模型学得稍慢,rank=32 时训练时间明显增加且没有看到效果提升,最终定在 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 散度衡量两个概率分布的差异。第二部分是让学生模型对真实标签做常规的交叉熵损失,防止学生模型完全被教师模型的错误带偏。这两部分的权重我最初设定为 7:3,后来发现偏重软标签时学生模型学到的“风格”更多,偏重真实标签时“事实正确性”更稳,最后调成 6:4。
温度 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 做对比来理解 DPO:RLHF 要训练一个 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.7:0.3 加权平均。这个方法虽然是土办法,但在小模型场景下经常能把倒退拉回来一些,代价只是多几次实验和一点点推理代码。
6.4 部署与推理阶段的小细节
最后提一下部署。300M 的模型虽然小,但推理框架的选择依然会影响实际效果。我直接使用 ONNX Runtime 导出模型,顺手做 INT8 量化,显存占用又降了一截。量化的损失通常不大,但要注意 tokenizer 部分也必须对齐,很多坑都出在导出后 tokenizer 和模型不一致上。
如果你要把模型嵌入到现有服务里,我强烈建议做一个“兜底策略”:当模型输出的置信度很低时,不要硬答,可以返回“需要更多信息”或触发旧规则逻辑。小模型的自信心和能力并不总成正比,兜底策略往往能提高整体用户体验,而不是把所有压力都压在模型上。
以我个人经验来说,训练一个小语言模型最核心的不是某一步有多高大上,而是每一步的输入质量是否配得上训练成本。Xihe 这个项目至今还在迭代中,后续我还会尝试更小更快的 checkpoint、更复杂的偏好数据构造,以及把整个流水线自动化。如果你也在做类似的实验,建议先从一到两个阶段跑通,再逐步往上加,别一上来就想着 Ablations 全做,那样数据、算力和时间都会被吃得很紧。