首页 / 资讯中心 / 文章详情

WAM模型训练策略全解析:从数据配比到后训练对齐

WAM模型训练策略全解析:从数据配比到后训练对齐 ★ FEATURED ARTICLE
最近整理手头项目时我把这两年围绕 WAM 类模型训练的相关工作重新过了一遍前后翻了近 300 篇论文、技术报告和开源仓库的 issue越看越觉得这类模型在训练策略上其实有很强的共性。WAM 指的是以“世界建模 动作决策”为核心目标的模型它跟纯粹的对话模型不一样不只要理解文本还要从多模态数据里抽取状态、预测变化、给出动作所以它的数据配比、预训练目标和后训练路径都不能直接照搬通用大模型的套路。这篇文章我把调研里反复出现的结论和踩过的坑统一整理出来按数据、预训练、后训练三个阶段展开最后附上实操中常见问题的排查方法。适合正在训练或者计划训练类似模型的同学参考不管你是刚入门还是已经跑过几版实验应该都能从中找到有直接价值的信息。1. 内容整体设计与思路拆解1.1 WAM 训练策略的分析框架WAM 这类模型有个很明显的特征它同时承担了“理解世界”和“执行动作”两种任务。这意味着训练过程不能只盯着语言模型的困惑度也不能只看某个 benchmark 的准确率而要建立一套从数据到训练目标再到评估指标的完整链路。我把近 300 篇工作里的训练策略拆成三层来看。最底层是数据层解决“模型看到什么”的问题包括数据来源、清洗、配比、增强和去重。中间是预训练层解决“模型学会什么”的问题包括训练目标设计、课程安排、稳定性控制和超参选择。最上层是后训练层解决“模型如何使用能力”的问题包括指令微调、偏好对齐和领域适配。这三层彼此耦合数据配比会直接影响预训练的收敛速度预训练的质量又决定了后训练的天花板所以不能分开孤立地调。调研里有一个很关键的观察很多 WAM 项目早期效果差问题不在模型结构而在数据配比失衡。典型情况是动作数据太少、文本数据太多模型变成了一个“话痨”对环境的感知和动作输出非常弱。这个现象在多个工作里都独立出现说明数据策略的优先级应该排在模型结构之前。1.2 近 300 篇调研的结论主线把这些工作通读下来我发现四条反复出现的主线。第一条主线是“数据质量 数据规模”。几篇有影响力的工作用 1B token 级别的干净数据超过了别人 10B token 的混乱数据关键在于过滤、去重和配比。第二条主线是“预训练要分阶段走”不能用一个统一目标从头训到尾先训基础能力、再训领域知识、最后训动作决策这种方式远比一次性混合训练稳定。第三条主线是“后训练必须解决遗忘问题”微调做多了模型会丢掉预训练阶段学到的世界知识所以要么用低学习率要么用模型合并要么引入回放数据。第四条主线是“评估指标必须跟着数据走”如果数据里只有少量目标场景样本评估结果几乎反映不出真实能力。这四条主线看起来简单但每一组背后都有大量对比实验支撑。后面几个章节我会把每一条展开讲清楚具体怎么落地以及有哪些坑是文献里没明说但实操时几乎必踩的。1.3 WAM 与通用大模型训练策略的差异对比通用大模型的训练流程WAM 有几个明显的差异点这些差异直接决定了策略选择。通用模型大多以文本为唯一模态数据管道相对成熟预处理工具链也很完善而 WAM 要处理视觉、文本、结构化状态数据可能还有时序信号数据格式差异大清洗难度成倍上升。通用模型的后训练通常以 SFT RLHF 为主对齐目标是“拟人化”和“有用无害”WAM 的对齐目标更复杂既要让模型遵守指令又要让模型真实理解物理世界或业务世界的约束。调研里有工作把后训练拆成两条并行分支一条负责语言指令跟随一条负责动作输出校准在模型内部通过路由机制合并。这种结构在通用模型里很少见但在 WAM 里效果很好。还有一点非常现实通用大模型可以用公开的网页语料做预训练WAM 需要的场景化数据大多来自行业内部流通性很差。所以很多 WAM 项目实际上是从数据采集和标注这一步就开始自己搭管道整个项目的重心往往不在模型结构上而在数据工程上。2. 数据策略从规模到质量的完整链路2.1 数据来源与采集行业数据才是主战场通用模型可以靠爬虫解决大部分数据问题但 WAM 不行。我调研的近 300 篇工作里凡是效果好的项目基本都有稳定的行业数据来源比如设备运行日志、传感器时序、业务数据库快照、操作轨迹记录等。这就引出一个很现实的建议如果你准备训练 WAM第一件事不是搭模型而是找到能持续产出高质量数据的业务系统。选数据来源时我建议优先关注“带标注的轨迹数据”。所谓轨迹数据就是“状态 动作 新状态”的三元组序列比如一个机器人从 A 点移动到 B 点过程中每一步的传感器读数、执行的动作、产生的反馈都被记录下来。这种数据直接对应 WAM 的核心能力远比单纯堆文本或堆图片有价值。很多项目组容易忽略这一点跑去下载一堆通用数据集结果模型学了一堆常识却不会做最基本的决策任务。我在实操中会常用到一些开放数据集做补充比如自动驾驶领域的 BDD100K、航拍目标检测的 HRSC2016、机械健康管理的 PHM2012它们能提供真实世界的视觉和时序样本做一些预训练阶段的能力预热是够用的。但请注意这些数据集只能作为“补充剂”不能作为“主食”主线数据一定得是目标场景的真实采样。2.2 数据清洗与过滤三类脏数据必须处理干净数据清洗这一节我要多写几句因为这是调研里出现频率最高的失败原因。所谓“脏数据”不只是格式错误或乱码更隐蔽的是以下几种情况。第一类是“重复数据”。尤其是从同一个业务系统里多次导出的数据很多在内容上是高度相似的只是时间戳不同。如果不做去重模型会对高频重复内容过拟合表现为对常规输入很流畅、对罕见输入非常差。常用的做法是 MinHash 加 LSH先对文本和图片特征做指纹计算再按相似度阈值合并这里阈值建议从 0.85 开始调太低了去不掉相似样本太高了又会误删多样性样本。第二类是“低质量样本”。对于文本部分可以用困惑度过滤也就是用一个小型的通用语言模型计算每句话的困惑度把明显偏离正常表达的片段丢掉。对于图片和时序部分可以用规则过滤比如图像分辨率过低、时序数据缺失率过高的样本直接移除。调研数据表明加入这一步之后下游任务准确率往往能提升 3 到 5 个百分点。第三类是“标签噪声”。WAM 的很多数据依赖人工标注或半自动标注出错率天然不低。我的经验是不要只靠人去复查可以训练一个简单的分类模型把标注结果做交叉验证把置信度低的样本抽出来人工二次确认。这一步耗费时间但能显著提升后训练阶段的对齐效果。2.3 数据配比与混合采样黄金比例到底存不存在数据配比是调研中讨论最热烈的话题之一但结论并不统一。不过有几个规律是跨项目成立的我在这里整理一下。第一预处理后的文本数据建议控制在总数据量的 40% 到 60%。WAM 不能没有文本能力因为指令理解依赖它但如果文本占比太高模型会把大量容量花在语言生成上动作决策能力会被挤占。第二动作轨迹数据至少要占 20%这部分是 WAM 区别于普通语言模型的核心能力来源比例再低模型的决策能力就会明显下降。第三视觉和时序状态数据各占 10% 左右具体比例取决于你的场景是图像密集型还是传感密集型。调研里有两个极端案例值得警惕。一个项目把文本数据堆到 90%结果模型在很多语言任务上表现不错但在真实环境测试时几乎不会输出有效动作。另一个项目把动作数据推高到 50%结果模型的动作输出很丰富却无法理解复杂的自然语言指令导致整个系统无法实用化。所以配比不能走极端要按“能力带宽”来分配。实际操作时我常用一个两阶段混合采样法先按比例确定每个数据源的量然后在采样时加入温度参数来控制多样性。温度越高越倾向于抽取小概率样本温度越低越集中在高频样本。预训练初期建议把温度调高一点让模型先见足够多样的数据训到中后期再把温度调低让模型在重点数据上做精调。2.4 数据增强与合成数据解决稀缺场景的可靠办法WAM 项目最痛苦的问题往往不是数据不够多而是目标场景的数据太少。比如你要做一个设备故障诊断模型正常运行的数据一大堆故障状态的数据却屈指可数。这时候靠采集是不现实的必须靠数据增强和合成数据来补。对于时序数据我常用的增强手段包括噪声注入、时间缩放、幅度扰动和通道屏蔽。噪声注入能让模型对传感器抖动更鲁棒时间缩放能提升模型对速度变化的适应力幅度扰动则有助于训练模型在量纲不一致时保持稳定。这些手段计算量不大效果却很明显。对于视觉数据基础的做法是随机裁剪、翻转、色彩抖动等再进一步可以用图像混合mixup或 CutMix 做样本插值。调研里有一个值得参考的结果在目标检测相关的 WAM 实验中使用 CutMix 后模型对遮挡目标的检测能力提升了约 7%这个方法几乎零成本强烈建议纳入标准流程。合成数据要复杂一些但解决的是真实数据覆盖不了的问题。你可以用模拟器生成大量带标注的轨迹数据也可以在已有数据上做状态插值。做合成数据时有一点必须注意一定要保留数据的物理约束比如一个机械臂的关节角度不能超过硬件限制否则模型学到的是无效的动作空间。调研里有个项目就在这上面吃过亏用仿真器生成了大量数据结果模型在真实设备上频繁撞到边界就是因为合成时没有加约束条件。3. 预训练策略稳定高效地把基础能力打牢3.1 预训练阶段划分先通用、再领域、后场景预训练阶段的划分是 WAM 训练策略里最体现功力的部分。通用大模型通常是一股脑把所有数据混在一起训但对 WAM 来说混合训练很容易导致能力互相干扰所以调研里主流做法是分阶段训练。第一阶段是通用能力预热。这个阶段用大规模通用语料和通用视觉数据让模型建立基本的文本理解和视觉感知能力。学习率可以相对高一些比如 3e-4 到 6e-4 这个区间因为模型参数处于快速收敛期数据量大采样密度高。第二阶段是领域知识注入。用行业数据、专业文档和场景相关的视觉数据让模型掌握具体领域的术语、状态模式和因果关系。这一阶段学习率要降下来我常用 1e-4 到 2e-4降幅太陡会导致前面学到的通用能力剧烈退化太缓又学不进去领域知识。第三阶段是场景能力精调。用轨迹数据和动作数据训练让模型把前面学到的知识转化为具体的决策能力。阶段之间怎么切换很多人直接用新数据覆盖旧数据其实更好的做法是用一个小的“衔接数据集”做过渡里面同时包含前一阶段和后一阶段的数据比例大概 3:7 到 5:5训几百步再切。这个做法在调研里被多次提到能明显降低训练 loss 的波动。3.2 训练目标设计不只是预测下一个 tokenWAM 的预训练目标不能只设为文本上的下一个 token 预测还要针对状态空间和动作空间设计专门的损失函数。调研中处理这个问题主要有三条策略路径。第一种是“多任务头联合训练”也就是在模型主干之上同时挂几个输出头一个负责文本理解一个负责状态预测一个负责动作生成。训练时对不同任务分配不同的损失权重文本部分权重低一些动作部分权重高一些让模型能区分优化方向。我在实际使用中会把状态预测的损失权重设为文本的 1.2 到 1.5 倍效果比默认的 1:1 更好。第二种是“辅助对比学习”。WAM 不仅要接受输入、给出输出还要理解不同状态之间的差异。所以训练过程中可以对状态表示做对比学习同一场景的不同扰动视为正样本不同场景的样本视为负样本让模型学会区分状态空间中的不同区域。这个方法在几个多模态 WAM 工作中效果显著状态表征的聚类效果有明显提升。第三种是“动态动作屏蔽”。训练时对动作序列做部分屏蔽要求模型从上下文和状态中推断被屏蔽的动作。这个思路借鉴了掩码语言模型的套路但应用在动作序列上能强制模型学到动作之间的依赖关系和时序逻辑。我当时在一个机械臂控制场景里试过加入动态动作屏蔽之后长序列动作预测的准确率提升了约 5%。3.3 超参数选择与计算逻辑超参数是预训练阶段最容易掉坑的地方。调研发现不同规模的 WAM 模型最优超参区间明显不同直接套用通用模型的参数设置往往效果很差。下面我按常见的三个规模档位给出参考区间。对于百亿参数以下的中小规模 WAM批次大小建议在 256 到 512 之间学习率用 3e-4 到 5e-4序列长度从 512 开始逐步增长到 2048。这个规模下模型容量有限过大的批次容易导致收敛不够充分但不能太小否则梯度噪声太大样本效率低。对于千亿级别的 WAM批次大小建议提升到 1024 甚至更高因为大模型对数据吞吐量的需求更高小批次反而很难收敛。学习率要降到 1e-4 到 2e-4同时要考虑增加梯度累积步数来缓解显存压力。还有个特别值得注意的参数是“上下文长度”。WAM 的动作决策往往依赖长距离的状态历史所以上下文长度不能太短。但把序列长度翻倍会导致计算量接近线性上涨对显存和训练时间影响很大。我的习惯做法是先把长度定在 1024 完成主要训练再用长度 2048 或 4096 做短阶段的续训这样既保证模型能处理长历史又不会让前期训练成本失控。3.4 loss 曲线波动与训练稳定性训练不稳定是 WAM 项目里最容易让团队崩溃的问题。常见表现是 loss 在正常下降过程中突然出现 spikes或者干脆在某一步之后开始发散。调研里有一个被反复验证的结论大部分训练不稳定问题都跟数据相关而不是模型结构出错。最典型的原因是“数据批次里出现了异常样本”。比如一批数据里混入了大量全黑图片或者全是重复文本梯度会突然异常表现为 loss spike。解决办法是建立数据采样时的质量监控每训练一定步数就在一个固定验证集上测一次 loss如果 loss 突然升高先回滚到之前的 checkpoint再去检查该 batch 对应的数据。梯度裁剪是另一道保险。WAM 的损失函数包含多个任务头梯度范数可能比较大不裁剪很容易让模型参数出现剧烈震荡。我通常把全局梯度范数限制在 1.0 以内如果是混合精度训练这个阈值可以放宽到 1.5 左右。warmup 步数也不能忽略。WAM 模型刚初始化时各任务头的输出分布差异非常大如果直接上全量学习率前期训练很容易崩溃。我建议 warmup 步数设为总步数的 1% 到 3%比如总步数是 10 万步就用 1000 到 3000 步线性升温。这个比例看起来小但在防止训练发散上非常有效。4. 后训练策略对齐、指令跟随与领域适配4.1 指令微调数据要怎么选、怎么配预训练做完之后模型具备的是“潜能力”要真正能用起来必须经过指令微调。但 WAM 的指令微调比通用模型更讲究因为在有限的微调数据里既要覆盖语言指令的理解又要覆盖动作输出的格式还要兼顾状态推理目标太多数据配比稍有不慎就顾此失彼。我常用的一个比例是文本指令样本占 50%状态-动作联调样本占 30%纯动作样本占 20%。这个比例下模型既不会变成“只会说话不会做”也不会变成“只会做但听不懂指令”。如果你手头资源有限只能做一个小规模的指令微调建议优先保证状态-动作联调样本因为那是 WAM 区别于聊天机器人的核心。指令微调的学习率要非常保守。通用模型的微调学习率在 1e-5 到 2e-5 是常见区间但 WAM 涉及多模态和多任务学习率高一点就可能让模型在预训练阶段学到的状态表征崩掉。我试过 5e-5结果验证集上的状态预测准确率直接掉了 8 个点换回 1e-5 之后效果才恢复正常。微调数据集的覆盖度也很关键。如果只在单一场景的数据上微调模型会变成“偏科生”。所以我在构建微调集时会刻意加入边界情况比如指令里出现未见过的说法、状态数据里出现异常值、动作目标与常规逻辑不一致等。这些边界样本可能只占微调集的一两成但恰恰决定了模型在真实环境里稳不稳。4.2 偏好对齐与模型合并别让后训练把预训练吃掉关于偏好对齐调研里头有几个反复被讨论的问题什么时候该用 RLHF什么时候用 DPO怎么避免对齐过程把预训练能力学没了。我的结论是如果偏好数据量少于几万条DPO 往往是更稳妥的选择。RLHF 需要先训练奖励模型再做强化学习采样整个过程对超参和数据质量极为敏感偏好数据不够时会严重过拟合。DPO 直接把偏好数据转化为隐式的奖励信号实现简单、稳定性好在小型团队的项目里是非常实用的方案。如果你确实需要 RLHF至少保证偏好数据里有足够多的“对比对”也就是同一输入下好坏两种输出的成对数据。但无论用 DPO 还是 RLHF都逃不掉“能力遗忘”的问题。模型在做偏好对齐时会逐渐放弃预训练阶段学会的一些低频知识表现就是对常规场景很好对冷门场景明显变差。解决这个问题有三个常用手段。一是加回放数据。在对齐训练的每个 batch 里混入部分预训练阶段的代表性样本比例控制在 5% 到 15%这能有效维持模型的通用能力。二是用低学习率 早停。监控一个包含预训练阶段任务的验证集当这个集合上的指标开始下滑时及时停止训练。三是模型合并。把偏好对齐后的模型和预训练模型在参数层面做加权平均这个做法的实现成本最低效果在某些场景下反而更好适合作为后备手段。4.3 领域适配少样本场景下的有效策略WAM 项目最常遇到的现实问题是团队拿到一个全新的领域手头只有几千条甚至几百条样本该怎么办调研里比较可靠的路径是“先检索、再生成、后筛选”的三步法。第一步把预训练阶段用到的领域数据做成索引库针对新的领域样本做相似度检索找出与这个领域最接近的历史样本作为微调数据的补充。第二步用已有的 WAM 模型结合少量新样本做数据扩展也就是让模型尝试生成新的描述或动作序列人工审查后挑出合理的。第三步根据扩展后的数据集进行小规模微调同时加入数据增强策略比如对时序样本做噪声扰动对文本样本做同义替换。这个方法虽然听起来不复杂但在调研中有多个项目靠它把冷启动的场景准确率从 40% 级别提升到了 70% 以上。核心在于把检索得到的历史知识作为先验注入让模型受新数据影响时有一个稳定的底座。4.4 评估策略后训练效果怎么看不被带偏后训练阶段最容易出现的误判是只看目标 benchmark 的分数不看模型综合能力。WAM 因为承担的任务多样单一指标根本无法反映真实水平。我的做法是建立“三级评估体系”。第一级是目标场景指标比如动作执行成功率、状态预测准确率这是最直接的业务指标。第二级是通用能力回测拿预训练阶段用过的通用任务集做回归测试看后训练是否损伤了基础能力。第三级是鲁棒性测试给输入加噪声、改变指令表述、打乱状态数据的顺序观察模型的输出是否稳定。在评估过程中我发现一个非常有价值的统计方法在目标场景指标之外额外记录模型输出的“格式合规率”。很多失败案例不是模型理解错了而是输出格式不对导致下游系统无法解析。把格式合规率放到评估体系里之后很多模型的真实可用度要打七折到八折。这个指标你可能觉得不值一提但在真实业务里是最决定上线成败的因素。5. 实操复盘与常见问题排查5.1 我的一次完整训练流程复盘纸上谈兵聊完了我拿一个实际项目复盘整个流程。这个项目要训练一个用于设备状态诊断和操作建议的 WAM输入是传感器时序和操作日志输出是诊断结论和推荐动作。数据方面我们从业务系统里拿到了大概 20 万条运行日志和 5 万条操作轨迹加上 PHM2012 和部分增强数据总样本量折算后相当于约 1.5B token。清洗阶段去掉了大约 18% 的重复样本和 6% 的低质量样本。配比上文本相关数据占 45%轨迹数据占 30%视觉和时序状态数据占 25%。预训练分两个阶段第一阶段用通用语料加公开数据集预热第二阶段用业务数据精调总步数 8 万步学习率从 5e-4 线性降到 1e-4。后训练做了 SFT 微调和 DPO 对齐SFT 数据约 8000 条DPO 偏好对约 3000 对。最后的结果目标场景的准确率从零样本的 52% 提升到 84%通用能力回测指标只下降了不到两个点。这个结果验证了前面说的几个核心原则优先保质量数据、阶段式训练、后训练用回放防遗忘。复盘时最值得执行的改进是缩短了第一阶段训练。我们原本规划通用预热跑 60% 的步数但实际跑到 35% 时通用能力已经达标省下来的算力全部投到了第二阶段领域精调上业务指标比原方案涨了 4 个点。所以别把阶段划分当成死规则多用验证集做动态判断。5.2 问题速查表与解决路径问题现象可能原因解决方法模型能对话但不会输出有效动作动作轨迹数据占比过低提高轨迹数据配比补充真实动作样本预训练 loss 频繁出现 spikes数据批次中混入异常样本检查对应 batch 数据回滚 checkpoint 后重训后训练后通用能力明显下降微调时未加回放数据或学习率过高混入 5%-15% 预训练代表样本降低学习率长序列状态理解差上下文长度不足用更长的序列做短阶段续训模型对未见过的指令没有反应指令微调数据覆盖度不足增加边界情况和同义改写样本输出格式不合规未将格式约束加入微调数据在训练和评估中同时加入格式合规率指标上面这张表里的大部分问题我都实际遇到过尤其是第一行“能说不会做”的现象出现频率极高很多团队一开始都以为是模型结构不行其实解法往往非常朴素——多喂点轨迹数据。还有一个容易被忽略的是数据和时间戳的问题。WAM 处理的多是带时间属性的状态数据如果数据被错误地混洗时间顺序被打乱模型学到的因果关系会完全错掉。调研里至少有三个项目在数据处理时踩过这个坑症状表现为模型对短序列的预测效果尚可对涉及长期依赖的任务几乎失灵。排查方法是抽几条样本人工核对时间顺序和原始日志的一致性。5.3 算力紧张时的策略取舍聊到实操绕不开算力预算这个现实约束。近 300 篇调研里同样有大量相关讨论核心结论是算力不足时优先保数据质量和后训练其次才是预训练规模。具体来讲如果算力只能支撑一个 7B 参数的模型做 5 万步训练我建议不要把步数硬撑到 10 万。更有效的方式是用 3 万步做分阶段预训练把省下来的算力用来做两轮 SFT 和一轮 DPO。这样模型在单一指标上可能不如硬撑 10 万步的版本但整体可用度和鲁棒性明显更好。另外一个省算力的小技巧是冻结部分层做后训练。WAM 的不同层对能力的分工不同底层更多编码通用感知信息高层更多编码语义和动作决策信息。后训练时可以冻结前 30% 到 40% 的层只更新后面的层参数量少了之后不仅训练更快遗忘问题也会缓解。我实测下来这个操作在目标场景指标上几乎无损失但训练时间和显存占用都有不少下降。如果你连微调的显存都很紧张可以考虑 LoRA 这类参数高效微调方法。在 WAM 上做 LoRA 时建议把 LoRA rank 设为 32 到 64 之间不要为了省显存把 rank 降到 8 以下。调研结果显示rank 过低时动作输出头学不到足够复杂的映射关系效果衰减明显。5.4 下一步可以从哪里扩展最后说一个我最近在探索的方向把 WAM 的预训练和后训练流程做成可配置的流水线让不同场景的项目可以共用一套数据管道和训练脚本。这个思路的出发点很现实因为 WAM 项目之间的差异往往不在模型结构而在数据格式和评估逻辑。如果把这些抽象出来做成通用组件新场景的启动时间可以从两个月压缩到两周左右。具体来说可以按“数据接入层、数据清洗层、训练调度层、评估层”四层拆解。数据接入层负责统一不同来源的格式清洗层负责去重、过滤和配比训练调度层负责阶段切换和超参管理评估层负责三级评估体系的自动化执行。这个架构在调研中已经有一些开源项目和论文尝试过但目前还没有形成统一标准值得跟进。我在实际使用中发现哪怕是简单的“把数据清洗规则做成可配置的 yaml 文件”这一步就能节省大量来回改代码的时间。特别是你要同时跑多个场景实验时规则可配置带来的收益非常大。更值得投入的是把评估层自动化每次训练完自动跑完整的三级评估并生成对比报告这会让你在做策略调整时有直接的判断依据而不是靠感觉拍板。
阅读完成 · 觉得有帮助?
咨询建站