基于多目标强化学习的LLM预训练全数据调度器
大语言模型(LLM)预训练中,训练数据的组成(来源多样性及其混合策略)是基石。在线数据混合(ODM)通过自适应调整数据配比提升效率,但现有方法局限于单一优化视角,忽略了复杂预训练对多维度动态数据组合的需求。 本文提出Holistic Data Scheduler(HDS),一种新型在线数据混合框架。HDS将数据调度建模为连续控制空间中的强化学习问题,采用Soft Actor-Critic(SAC)算法以在稳定性和样本效率上探索高维策略空间。其核心是多目标整体奖励函数,融合三个维度:1)数据驱动奖励衡量质量;2)损失驱动奖励捕获跨领域影响;3)模型驱动奖励基于权重范数。 在The Pile基准上,HDS达到最优方法的最低验证困惑度,训练迭代减少44%;在MMLU 0-shot任务中提升7.2%,并在其他基准上取得一致增益,验证了其提升训练效率与模型能力的双重优势。
论文精读
TL;DR 用多目标强化学习(SAC)动态调整预训练数据配比,以 44% 更少迭代达到同等困惑度,MMLU 提升 7.2%,兼顾训练效率与模型能力。
问题
问题背景
LLM 预训练的性能高度依赖训练数据的组成和动态混合策略,在线数据混合(Online Data Mixing, ODM) 能够在训练过程中自适应调整各数据源的比例,以提升训练效率和最终模型能力。
现有方法的局限
当前 ODM 方法普遍基于单一优化视角:
- 损失驱动类方法(如 DoGE)仅关注域间损失变化,容易忽略数据多样性和模型稳定性,导致域权重震荡或过拟合;
- 数据驱动类方法(如静态启发式或基于多样性的调度)缺乏对模型动态反馈的感知,难以在训练后期有效倾斜高价值数据;
- 现有工作将数据混合视为离散选择或简单加权,无法在连续控制空间中精细调节比例,且未显式建模多目标冲突。
这导致训练效率未达上限,且最终模型能力(如推理、知识覆盖)提升有限。
问题为何重要且困难
LLM 预训练是一个高维、非平稳的复杂系统:数据混合会同时影响模型参数更新方向、损失空间分布和泛化能力。设计一个能同时兼顾以下三个维度的调度器极具挑战:
- 数据质量:保证 token 多样性和语义丰富度;
- 跨域影响:量化不同数据源之间的 loss 信号对齐;
- 模型状态:通过参数范数等指标监控训练稳定性。
业界对降低预训练成本、提升模型素能的需求迫切,据估计将 175B 模型预训练效率提升 10% 可节省数百万美元计算费用,因此该问题的任何进展都有巨大工程价值。
行业类比
类似多目标强化学习在推荐系统中平衡多样性、准确性和用户留存,但 HDS 将这一思路迁移到训练数据调度,在连续动作空间中动态决策数据混合比例,直接优化预训练累积收益。
核心洞察
- **将在线数据混合从单目标启发式规则升级为多目标强化学习控制问题** 现有 ODM 方法往往基于一个孤立指标(如单域困惑度或梯度相似度)调整混合比例,忽视预训练数据动态组成的多维特性。HDS 首次将调度形式化为连续空间 RL 任务,由 SAC 智能体根据训练状态(各域损失、模型权重范数、已消耗数据比例)实时输出连续混合权重,并在奖励函数中同时编码**数据质量、域间影响、模型稳定性**三个目标,使策略能够探索更灵活、整体更优的数据配比,突破单目标贪心优化的短期偏差。
方法
方法概览
HDS 将 LLM 预训练的在线数据混合 (ODM) 问题形式化为一个连续控制空间下的强化学习任务,核心组件包括:
- 输入:从当前训练状态中提取的多维信号(各域验证损失、词元多样性、模型权重范数)。
- 关键模块:SAC Agent 作为策略网络,输出各数据源的采样权重(连续动作);Holistic Reward Function 驱动 Agent 学习最优混合策略。
- 输出:每个训练步的域采样概率分布,指导数据加载器按比例混合 batch。
状态、动作与策略
- 状态空间:包含各域的近期损失值、词汇多样性指标,以及模型部分层的权重范数(默认取倒数第二层 Transformer block 的输出范数)。
- 动作空间:一个与数据域数量相同的连续向量,经 softmax 归一化后得到各域的采样概率。
- 策略:使用 Soft Actor-Critic (SAC) 算法,通过最大化累计奖励与策略熵的加权和,平衡探索与利用;SAC 的稳定性和样本效率适合高维连续动作空间探索。
Holistic Reward Function
奖励由三个子项线性加权组合,分别对应数据、模型、损失三个视角:
- Inter-Domain Influence Reward (
r_align):基于不同数据域上损失的互信息,评估域间正向或负向影响。它鼓励调度器优先选择那些能降低其他域损失的域,捕获跨域知识迁移。 - Scheduled Lexical Diversity Reward (
r_diversity):衡量实际采样数据的词元多样性(如去重 token 比例),防止模型因过度集中于单一域而导致表征坍缩。 - Model Stability Reward (
r_stability):根据模型参数范数的变化(如权重矩阵的 Frobenius 范数增长率)设计,惩罚剧烈参数变化,避免灾难性遗忘。
最终奖励为三者的加权和:R = α·r_align + β·r_diversity + γ·r_stability,超参数通过消融确定(论文中默认 α=1.0, β=0.5, γ=0.1)。
训练循环
每一步,SAC Agent 接收当前状态,输出采样权重;数据加载器据此构造 batch;LLM 梯度更新后,收集新状态并计算多目标奖励;SAC 使用 replay buffer 进行软策略迭代更新。LLM 与 SAC Agent 异步更新,SAC 更新频率低于 LLM,以降低额外计算开销。
与同类方法的差异
不同于 DoReMi、DRO 等仅基于单一指标(如域损失或梯度对齐)的在线调整方法,HDS 首次在连续控制 RL 框架下联合考虑数据多样性、模型稳定性和跨域影响,其多目标奖励设计使调度器能捕捉复杂预训练动态的更多维信号。
实验
实验设计
训练数据使用 The Pile 基准,涵盖多种规模的 LLM 预训练。HDS 将数据调度建模为连续控制空间的强化学习问题,采用 Soft Actor-Critic (SAC) 算法动态调整各数据源混合比例。奖励函数融合三个维度:数据驱动质量奖励(词汇多样性)、损失驱动域间影响奖励(损失对齐)和模型驱动稳定性奖励(权重范数)。对比基线包括静态混合策略及其他在线数据混合方法,评估指标包括验证困惑度及下游任务(如 MMLU 零样本)性能。
关键发现
- HDS 以比次优基线少 44% 的训练迭代达到相同验证困惑度,显著提升训练效率
- 在 MMLU 0-shot 上实现 7.2% 的绝对提升,且在其他基准(如 LAMBADA)上观察到一致增益
- 消融实验表明,三项奖励组件协同作用至关重要,缺一不可
- SAC 的高维策略探索稳定性使得数据配比在训练过程中平滑演化,避免剧烈波动
与基线对比的深度解读
现有在线数据混合方法多依赖单一优化视角(如仅最小化域损失或最大化数据量),忽略了预训练数据的多维动态特性。HDS 通过多目标奖励函数同时捕捉数据质量、域间影响和模型稳定性,使调度策略更贴近真实训练需求。44% 的迭代缩减意味着在同等算力下可提前达到目标性能,或节省大量计算资源。MMLU 提升表明,更智能的数据组合不仅能加速收敛,还能改善模型的泛化能力和知识覆盖。该框架为大规模预训练的数据工程提供了新范式,其强化学习驱动的设计易于扩展到更多数据源和更细粒度的调度任务。
行业影响
落地场景
HDS 可嵌入任何大规模语言模型预训练流程,尤其适用于需要从多源异构数据(如网页、书籍、代码、学术论文)中学习的场景。内容平台、云服务商、企业 AI 部门在训练自有基础模型时,可直接用 HDS 替代静态数据配比或手动启发式调整。典型产品形态包括:
- 云厂商的 LLM 训练平台(如 Amazon SageMaker、Google Vertex AI)可将其作为内置调度策略;
- 企业内部 领域大模型训练管道(如金融、医疗、法律垂直模型),动态平衡通用知识与专业数据;
- 开源训练框架(如 Megatron-LM、Hugging Face Trainer)中的即插即用数据混合模块。
商业价值
HDS 从两个维度带来直接商业回报:
- 降本增效:同等模型质量下训练迭代减少 44%,显著缩短 GPU 占用时间,降低云成本或加速迭代周期。这对千亿参数模型训练尤为关键,可节省数百万美元计算预算。
- 性能提升:在 MMLU 等基准上 7.2% 的零样本提升,直接改善模型在开放域问答、代码生成、多语言等下游任务的表现,增强产品竞争力。 同时,其样本高效的特性减少了对大规模调参实验的依赖,进一步压缩研发成本。
与现有工作流的接口
HDS 以 RL agent 形式工作,不侵入模型结构,可与现有训练 stack 集成:
- 数据加载层:维护各域数据源的预处理管道,HDS 输出下一批次的域比例
w_t,控制采样权重; - SAC agent 更新:从训练循环中定期提取状态(
loss分布、权重范数、数据多样性指标),计算奖励后通过异步进程更新策略; - 监控与干预:提供实时的域混合权重仪表板,允许工程师根据观测手动覆写或设置约束。
具体落地用例
- 电商搜索与推荐底层模型训练:电商平台训练联合理解商品描述、用户评论、搜索查询的 LLM。HDS 可动态分配商品文本、对话式评论、多语言查询的比例,防止模型过度偏向某类数据,提升长尾 query 理解和冷门商品推荐的泛化性。
- 自动驾驶多模态基础模型预训练:自动驾驶公司需融合传感器描述文本、驾驶手册、路测日志自然语言标注等异构数据。HDS 的
ralign奖励可捕捉域间知识迁移信号,自动平衡安全关键场景数据与常规路况数据的混合,在不牺牲通用语言能力的前提下,强化模型对罕见边缘案例的理解。
局限
- 实验仅在 The Pile 和有限模型规模(最大 1.3B)上验证,尚未在更大规模模型(如 7B+)或更广泛的数据集(如 SlimPajama、Dolma)上证明泛化性。SAC 代理引入的额外计算开销(如状态特征提取、策略推理与环境交互)可能抵消部分训练迭代节省的收益,文中未详细量化该开销与总训练时间的权衡。
- 多目标奖励函数中的权重(λ₁、λ₂、λ₃)需人工设定,且对性能敏感;文中的超参数优化(附录)基于固定搜索,缺乏自动平衡机制。不同训练阶段或不同模型架构可能需要重新调参,限制了该方法在实际生产环境中的即插即用性。
- HDS 的动态权重调整虽在困惑度和下游任务上有提升,但其决策过程缺乏可解释性,无法清晰说明为何某个域的采样比例应增加或减少,这对于数据分析导向的团队可能构成采用障碍。