搞了大半年,几乎把精力都砸在这套 WAM 模型上。WAM 是我们内部对一类以“通用任务理解与执行”为目标的预训练模型的代号,当时立项的原因很朴素:线上业务里到处是零散的任务型需求,每个场景单独训练一个小模型成本太高,于是想做一个能统一支撑的底座模型。为了不走弯路,前前后后翻了近 300 篇相关工作调研,从数据清洗、预训练调度到后训练对齐,把公开论文、开源实现和团队踩过的坑全部对照了一遍。这篇就把训练策略里最值得说的部分整理出来,重点讲清楚数据、预训练、后训练三条线怎么配合,以及那些论文里不会明说但实测很好用的细节。
如果你正要训练类似的通用模型,或者正在纠结“到底先扩数据还是先调对齐”,这篇应该能给你省下不少试错时间。我的原则是尽量说人话,所有结论都尽量落到可执行的参数和步骤上,能用表格的地方不废话。
1. 调研范围与整体设计思路拆解
1.1 近300篇调研都在解决什么问题
这 300 篇调研不是一次性看完的,而是按“数据工程—预训练架构—后训练对齐—工程稳定性”四条线分批推进。最初两周主要看数据相关的,包括大模型数据配比、质量过滤、去重、合成数据;中间一个月集中看预训练,重点对比了几种学习率调度和阶段划分方法;后面一段时间全在看后训练,包括 SFT 数据构造、DPO、模型合并。
看完之后一个很明显的感受是:大家讨论得最多的往往不是模型结构,而是数据策略。很多论文的实验结果差异,本质上就是数据配比和质量控制方式的差异。比如同样规模的模型,有的团队用 1T token 训出很好用的模型,有的团队堆到 3T token 效果却一般,差距基本都在数据上。
调研过程中我也把结论做成了几张对照表,这里列一下核心脉络:
| 主题 | 代表方向 | 对我们的启发 |
|---|---|---|
| 数据配比 | 文本、代码、多模态混合比例 | 代码数据远比想象中重要,必须单独保证 |
| 数据清洗 | 规则过滤、分类器过滤、嵌入去重 | 质量优先于数量,清洗要分多轮 |
| 预训练调度 | WSD、余弦、线性三个阶段 | WSD 在收敛速度和效果上表现更稳 |
| 后训练对齐 | SFT、DPO、RLHF、模型合并 | DPO 性价比最高,适合资源有限团队 |
1.2 训练策略的整体框架:数据、预训练、后训练三者如何协同
想把 WAM 这类模型做好,最重要的一点是先想清楚数据、预训练、后训练各管什么。数据决定能力的上限,预训练决定基础能力的下限,后训练决定模型最终交付时的表现形态。
打个比方:数据是食材,预训练是厨师基本功,后训练是摆盘和调味。食材好但厨师不会做,浪费;厨师厉害但食材不行,上限就在那。后训练只能放大已有能力,不能凭空创造新能力,这一点我在后面的实验里反复验证过。比如尝试仅靠 SFT 让模型学会逻辑推理,效果极差,推理能力必须在预训练阶段就已经形成雏形,SFT 只是把这种能力引导到指定输出格式上。
所以在设计整体框架时,我们定了三条原则:
- 数据建设先行,读作“先有干净的大锅饭,再谈小灶”。
- 预训练阶段把基础能力做扎实,不急着对齐。
- 后训练阶段只做引导和适配,不指望它改命。
这三条原则贯穿了后续所有实验,也直接影响了我下面要写的每一步细节。
2. 数据篇:喂给WAM模型的“食材”怎么选、怎么洗
2.1 数据来源与配比设计
WAM 一开始定位就不是纯文本模型,它还要接图像输入和结构化数据,所以数据来源天然比单纯的语言模型要复杂。我们的数据池大致分成四类:通用文本、代码、图文对、结构化数据。
通用文本来自网页、书籍、论文摘要、维基百科等,是基础语言能力的来源。代码数据来自开源代码库和代码问答社区,这部分对我们的任务理解能力提升非常显著。图文对我们用了公开数据集做了清洗和重采样,主要包括自动驾驶、车牌检测、遥感影像这些场景的公开数据,比如 CCPD、BDD100、HRSC2016。结构化数据则来自业务侧的表格、数据库导出和传感器时序记录,像 PHM2012 这类公开的时序数据集也会放进预训练数据里做辅助验证。
配比上我们试过很多组,最终稳定在这个比例:
- 通用文本:约 50%
- 代码数据:约 25%
- 图文多模态:约 15%
- 结构化与表格数据:约 10%
为什么代码要占到四分之一?因为代码天然带有严格的逻辑结构、符号推理和长距离依赖,这对模型学习“逐步计算”和“指令跟随”非常有帮助。很多论文也都验证过,代码数据能显著提升模型在数学和逻辑任务上的表现。我们自己在评测中也发现,只要代码配比低于 15%,模型在复杂指令上的表现就会明显下降。
2.2 清洗、去重与质量过滤的实操要点
数据配比定好了,真正耗时间的是清洗。我们第一版数据管道跑完,发现文本里有大量乱码、HTML 残留、无意义重复和敏感信息。这些脏数据不清理,后期的对齐实验会非常痛苦,因为模型会在无关特征上过拟合。
清洗分四步走:
- 编码修复和格式归一化:统一转成 UTF-8,过滤掉乱码字符。
- 规则过滤:去掉过长无换行的文本、纯标点文本、广告模板、验证码类文本。
- 去重:先做精确哈希去重,再做 MinHash 近重复去重,阈值设在 0.8 左右。
- 质量过滤:训练一个小分类器,基于困惑度打分,过滤掉低质量句子。这个分类器也可以用现成的预训练模型改造,但实测自己训一个轻量二分类模型更可控。
去重这一步特别提醒一下:不要只做文档级去重,还要做段落级去重。模型反复看到同一段话,会不自觉地复制训练语料里的原句,这在问答场景里非常致命。我们第一轮只做了文档级去重,结果后训练阶段发现模型经常原封不动吐出语料里的整段话,后来补了段落级去重才改善。
清洗过程中我也做了不少数据可视化,把每类数据的长度分布、困惑度分布、重复率画出来看。画图这个动作看上去很“不技术”,但实际能省下大量时间。比如有一批数据长度曲线出现双峰,明显是两类来源混在一起,需要重新分流。不画图靠肉眼统计很难发现。
2.3 数据增强与合成数据
当真实数据不够用的时候,再考虑数据增强和合成数据。WAM 的多模态部分我们用了常用的增强方式:随机裁剪、颜色扰动、旋转、翻转。这些对视觉任务有效,但要注意不能改变语义,比如车牌识别数据不能做水平翻转,否则字符顺序就反了。
文本增强我们做得不多,主要靠回译和掩码重建。回译适合扩充对话场景数据,但成本高,且翻译会引入语义漂移。掩码重建更适合做预训练阶段的辅助任务,不适合直接作为 SFT 数据。
合成数据我们重点用在两个地方:一是构造指令数据,二是平衡类别分布。比如某些任务在真实数据里占比只有 0.1%,模型总是学不好,我们就用模板加规则生成一批合成样本,把占比拉到 1%。但要注意,合成数据过多会让模型产生模式单调的问题,我们的经验是合成数据不要超过数据池总量的 5%,超过后模型输出会明显变得模板化。
3. 预训练篇:从零开始教WAM模型“说人话”
3.1 预训练的目标与阶段划分
预训练的目标不是让模型立刻变成领域专家,而是让它具备通用的语言理解、推理和生成能力。我们采用的方案是两阶段预训练。
第一阶段是通用预训练,目标是让模型掌握语言规律和世界知识。这个阶段用的数据就是前面说的四类混合数据,上下文长度从 4k 逐步扩展到 8k、16k、32k。第二阶段是高权重精调,相当于把最后数据的质量权重提上去,让模型在收尾阶段见到更多高质量语料。这个思路有点像期末冲刺:平时多门课一起学,最后半个月集中刷重点题。
两阶段之间我们用了连续训练,不重置优化器状态,只是调整数据采样权重。这样过渡比完全中断再重新训练稳定得多,loss 不会出现明显回升。
3.2 关键超参:学习率、Batch Size、上下文长度
预训练的超参我直接说实测配置,都是跑过大量对比实验后沉淀下来的。
学习率调度方案我们最终选择了 WSD(Warmup-Stable-Decay)而不是传统余弦。WSD 的思路是大部分训练步数保持稳定学习率,最后 10% 到 20% 的步数内衰减到最低点。这样做的好处是可以在稳定阶段随时拉出来一个 checkpoint 做评测,效果都差不多,最后再统一做衰减,模型收敛得反而更干净。
具体参数参考:
| 超参 | 取值 | 说明 |
|---|---|---|
| 峰值学习率 | 5e-4 | 只在 warmup 结束时达到一次 |
| 最低学习率 | 1e-5 | 衰减阶段终点 |
| Warmup 步数 | 2% 总步数 | 太短容易震荡,太长浪费时间 |
| Batch Size | 动态增大 | 稳定阶段 4M tokens 左右 |
| 上下文长度 | 4k -> 32k 递增 | 前面用短上下文快速见多数据,后面用长上下文精调 |
Batch Size 这里多说一句。我们不是一直固定一个值,而是前 10k 步用较小的 batch 让训练稳定下来,之后逐步增大。增大 batch 的时候要同步调整学习率,经验公式是“学习率随 batch 大小的平方根缩放”,不调整的话 loss 曲线会出现明显台阶。
3.3 课程学习与多阶段预训练
课程学习的想法很自然:先学简单的,再学难的。但实践中要小心。我们在早期实验里尝试过先纯文本后代码的训练顺序,效果并不好,因为模型把文本基础打得过死之后,再引入代码时会产生灾难性遗忘。
后来改成了“并行混合 + 后期加权”的方式。前 60% 步数里四类数据按基础配比混合,让模型自然建立多维能力;60% 到 80% 步数逐步提高代码和结构化数据的采样权重;最后 20% 步数配合 WSD 衰减,同时加大高质量领域数据的比例。这个顺序实测比严格课程学习更稳定,也不容易出现遗忘问题。
另外补一个细节:多模态图文对数据不要在前 10% 步数就大量引入。模型文本基础还不稳的时候,图文对会造成干扰,表现为训练 loss 迟迟不下降。等文本 loss 进入平台期之后再加入图文对,整体训练过程会顺很多。
4. 后训练篇:让模型学会“干正事”
4.1 SFT阶段的任务设计与数据组织
预训练结束之后,模型其实已经具备很强的基础能力,但输出格式可能不是我们想要的。SFT 就是在这个阶段教它“回答问题要规范、要按指令走”。
我们的 SFT 数据量大约做了 30 万条指令样本,这个数量不算多,但覆盖了 20 多个任务类型。每一类任务我们都设计了明确的指令模板和输出格式要求。比如数据类任务要求模型先复述分析思路再给结论,表格任务要求模型只输出表格结构,代码任务要求模型给出可运行代码而不是伪代码。
构建 SFT 数据时最忌讳的是只改开头套话、中间内容不动。我们专门做了指令多样性检查,统计每条指令的 token 重合度,重合度太高的样本直接降低采样权重,否则模型会对某一种表达方式过拟合。
训练时 SFT 一般只需要 1 到 2 个 epoch。超过 2 个 epoch 模型非常容易过拟合训练集,表现是训练 loss 降到极低,但评测集分数反而下降。我们第一版 SFT 跑了个 3 epoch,结果模型在泛化任务上掉了好几个点,后来强制规定 SFT 不超过 2 epoch 才稳定下来。
4.2 偏好对齐与DPO实操
对齐阶段我们用的是 DPO 而不是 PPO。原因很简单:DPO 不需要单独训练奖励模型,资源占用小,实现也简单得多。在 WAM 这种已经跑完 SFT 的模型上,DPO 只要准备偏好对数据就行。
DPO 的偏好对数据我们构造了一万多条,每条包含 chosen(更优回答)和 rejected(更差回答)。这些数据一部分来自人工标注,一部分来自规则筛选。人工标注比较贵,我们只在高价值任务上做;规则筛选适合那些有明确客观标准的任务,比如代码能跑通就算 chosen,跑不通就算 rejected。
实际操作中 DPO 的关键参数是 beta,它控制模型对偏好差异的敏感度。beta 太大,模型几乎不会偏离 SFT 结果;beta 太小,模型会过度迎合偏好数据,输出变得不自然。我们调试下来的经验是 beta 取 0.1 到 0.3 之间。初始用 0.1 跑一轮,看生成样本的多样性,如果感觉模型说话太模板化,就调到 0.3 再试。
DPO 训练时还有一个常见问题:学习率要设得比 SFT 更低。我们 SFT 用的学习率是 1e-5 量级,DPO 只有 1e-6。因为 DPO 是在已经收敛的模型上微调,学习率大了极易破坏原有能力。
4.3 模型合并与多任务增强
后训练阶段我们实验了另一种非常实用的做法:参数合并。如果某个垂直场景有一批 SFT 数据,可以在通用 SFT 之后,用这批垂直数据单独微调一份模型,再把这份模型的参数和通用模型合并。
合并方式我们试过两种:一种是直接在权重空间做插值,公式是“新模型 = 基础模型 + 系数 × (垂直模型 - 基础模型)”,系数一般在 0.3 到 0.5 之间。另一种是拼接多个垂直模型做策略路由,让输入根据任务类型走不同的模型分支。后一种效果好一些,但推理成本高,所以我们大部场景还是用了权重插值。
权重插值的好处是能保留通用能力,同时强化垂直任务表现。集合了多组实验后,我们发现垂直数据量少的时候(几千条),权重插值效果远好于直接合并训练;垂直数据量过万后,合并训练略优。因此判断标准就是数据量。
5. 训练过程的工程细节与稳定性保障
5.1 训练监控与loss曲线解读
训练大模型最怕的不是跑得慢,而是跑着跑着 loss 开始异常,你还没发现。我们第一版训练就吃过亏,因为只看全局平均 loss,某个数据子集质量崩了根本看不出来。
后来我们改成按数据类别分别记录 loss。文本 loss、代码 loss、图文 loss、结构数据 loss 全都单独打点,每 500 步画一次曲线。一旦发现某一类 loss 突然上升,基本能锁定对应数据源出了问题。
正常的 loss 曲线应该是整体下降,波动幅度随步数增加逐步收窄。如果看到 loss 曲线出现周期性凸起,大概率是数据采样顺序问题,某个数据块整体质量偏低被周期性采到。这时候不要急着调学习率,先去查数据管道。
5.2 算力规划与checkpoint管理
WAM 这套模型我们跑在 8 卡到 32 卡不等的集群上,训练过程横跨几周。算力规划很现实的问题:显存不够怎么办?我们的经验是优先用梯度累积和混合精度,实在不够再做序列并行和 ZeRO 分片。
混合精度这块我们用的是 BF16,比 FP16 稳得多,训练过程中几乎不用处理溢出问题。FP16 在训练早期容易丢精度,后期容易出现梯度爆炸,BF16 能省掉大量调 loss scaler 的时间。
checkpoint 管理也很重要。我们每 2000 步存一个完整 checkpoint,每 500 步存一个仅优化器状态的轻量 checkpoint。完整 checkpoint 用于回退和评测,轻量 checkpoint 用于解算力崩溃时快速恢复。磁盘空间建议按模型参数量的 500 倍预留,否则很容易在训练中途因为磁盘满了而中断训练。
5.3 实战中的坑与排查手册
训练过程中踩过的坑太多了,挑几个典型问题列成表,方便直接对照:
| 现象 | 可能原因 | 排查方法 |
|---|---|---|
| 训练 loss 不下降 | 数据清洗不到位或 batch 太小 | 检查数据样本,先跑小规模实验 |
| loss 曲线台阶式跳变 | 增大 batch 后未同步调整学习率 | 按平方根缩放学习率 |
| 模型输出重复话 | 段落级去重不够 | 补充段落级 MinHash 去重 |
| SFT 后泛化下降 | epoch 过多或指令多样性不足 | 控制 epoch 1-2,做指令重合度统计 |
| DPO 后模型说话模板化 | beta 太小或学习率偏高 | beta 调到 0.3,学习率降到 1e-6 |
| 推理时多模态效果差 | 图文对数据引入过早 | 预训练前 10% 步数不加图文对 |
还有一个很隐蔽的坑:多机训练时不同机器的数据打乱种子必须一致,否则数据分布出现偏差,loss 会异常。我们因为这个问题排查了整整一天,最后发现是三号机的 seed 没设对。
另外,训练中途如果出现显存不足,第一反应不是硬调 batch,而是检查是否有 tensor 没有被释放。我们在日志里加了显存打点,专门盯峰值显存位置,定位了几处多余的中间变量缓存,省出来的显存足够把 batch 再加大 20%。
写在最后的实操体会
WAM 这套模型从调研到训练稳定,整个过程最大的体会是:不要试图一步到位,每一项策略都要允许单独做对比实验。数据配比、学习率、SFT epoch、DPO beta,这些参数之间是耦合的,但不代表不能拆开调。我们每次只改一个变量,记录完整结果,才最终沉淀出这套配置。如果你时间有限,我建议优先把数据清洗和段落级去重做好,这个收益最大;预训练调度可以抄作业用 WSD,稳定且省心;后训练阶段 SFT 控制在 2 epoch 以内,DPO 的 beta 从 0.1 往上试。按这个顺序来,你的 WAM 型项目大概率能少熬夜、少烧钱,拿到一个靠谱的模型。