CausalMix:把 SFT 数据配比当成因果推断问题
RegMix 用 512 个 1M 代理模型 + LightGBM 回归解决了预训练阶段的配比优化问题,但它的假设是静态的:拟合一次,给出一组全局最优配比。当底层数据池换了(域增减、分布偏移),整套回归需要重跑。到 SFT 阶段,这个问题更突出,因为 SFT 的数据池和域定义高频变动。
CausalMix( Data Mixture as Causal Inference for Language Model Training ,2026/07)把配比优化重新建模为因果推断问题。区别在于:RegMix 学的是一个全局映射 $T \to Y$(配比到性能),CausalMix 学的是一个条件因果效应 $\theta_0(X)$(在当前数据状态 $X$ 下,配比的边际回报)。通过 Double Machine Learning 把数据状态的混杂效应正交化掉,因果模型在迁移实验中不需要对新数据池重做代理实验。
在 tulu-3-sft-mixture 上用 512 个 Qwen2.5-0.5B 拟合因果模型后,外推到 800K 样本训 7B 模型,CausalMix-S 在 AvgDev 上达到 62.28(RegMix 60.14、DMO 60.35),CausalMix-A 在 AvgUns 上达到 49.09(RegMix 48.12、DMO 48.98);迁移到完全不同的 Qwen3-4B + AM-Thinking 长思维链数据上,CausalMix 平均分 66.66,超过 RegMix 61.40 和 DMO 63.47。
问题背景与动机
RegMix 的回归目标是一个验证集上的 loss:给一组配比 $T$,预测 Pile-CC 的 validation loss。这个设定在预训练阶段能工作,但到 SFT 阶段遇到两个困难:
- SFT 的评估指标是下游任务准确率而不是 loss,loss 和下游分数之间的对应关系没有预训练那么稳定
- SFT 的数据池高频变动(换个 instruction set、加一批 code 数据),RegMix 式的回归需要对每个新数据池重新跑 512 个代理
CausalMix 的动机是:如果能学到"在当前数据状态下,配比的因果边际回报",那即使数据池变了(反映为数据状态 $X$ 变了),因果模型只需要把新的 $X$ 代入就能给出新的最优配比,不需要重新做代理实验。
因果建模
论文把问题建模为:给每个代理实验 $(X_i, T_i, Y_i)$,其中 $X$ 是数据状态(具体为三个指标:HES 衡量推理复杂度、Normalized_Loss 衡量数据难度、Writing_Style 衡量文本质量),$T$ 是配比向量,$Y$ 是 Development set 上的下游任务平均分。目标是估计条件边际回报:
$$\mu(x, Z) \approx g(x) + \theta_0(x)^\top Z$$
其中 $Z = \log(T + \epsilon)$ 是 log-mixture 表示,$g(x)$ 是基线性能,$\theta_0(x) \in \mathbb{R}^K$ 是 CATE(条件平均处置效应),即"在数据状态 $x$ 下,对 log-配比 $Z$ 做微小扰动后性能如何变化"。
如果 $\theta_{0,k}(x) > 0$,说明在当前数据状态下增加第 $k$ 个域的比例能提升性能;如果 $< 0$,则存在负迁移。
DML 正交化
直接回归 $Y = f(X, T)$ 会把数据状态的混杂效应和配比的因果效应混在一起。论文用 Double Machine Learning(DML)做正交化:先用两个 LightGBM 分别估计 $\hat{m}(X) = E[Y|X]$ 和 $\hat{e}(X) = E[Z|X]$,然后用残差做因果效应估计:
$$\tilde{Y} = Y - \hat{m}(X), \quad \tilde{Z} = Z - \hat{e}(X)$$
$$\hat{\theta} = \arg\min_\theta \sum_i |\tilde{Y}_i - \theta(X_i)^\top \tilde{Z}_i|^2$$
这个 R-loss 目标不优化绝对分数的预测,而是估计"配比的残差变动如何解释性能的残差变动"。最终的 CATE learner 用 CausalForestDML 实现,即因果随机森林。

上图是完整流程。左侧从 tulu-3-sft-mixture 的五个域出发,采样 512 个子集作为 treatment $T$;中间用 OpenDataArena 的预计算分数提取三个 covariates $X$;右上用 Qwen2.5-0.5B 做 proxy 训练并在下游任务上评估得到 outcome $Y$;右下把 $(T, X, Y)$ 组成 meta-dataset 后做 DML 正交化,用 CausalForestDML 估计 CATE。
消融实验验证了正交化和数据状态 $X$ 的必要性:去掉正交化(w/o Orth.),模型退化为直接拼接 $(X, T)$ 做监督回归,AvgDev 降到 59.65,低于 RegMix 的 60.14,说明不做正交化还不如全局回归;去掉 $X$(w/o X)AvgDev 为 61.30,也低于完整 CausalMix-A 的 61.84。
从 CATE 到最终配比
拿到 $\hat{\theta}(X_{\text{tar}})$ 后,论文给了两种从边际回报转化为配比的方法:
解析法(CausalMix-A):直接取正的边际回报做归一化,$T_k \propto [\hat{\theta}_k(X_{\text{tar}})]_+$
搜索法(CausalMix-S):从 Dirichlet 分布采 100,000 个候选配比,用因果模型预测每个候选的得分,取 top-100 平均。这和 RegMix 的搜索策略一致,差别在于打分函数从回归模型换成因果模型。
实验里 CausalMix-S 在 unseen 集上通常稍优于 CausalMix-A,论文把这归因于 top-100 平均起到了平滑单点噪声的作用。
实验结果
数据用 tulu-3-sft-mixture,分五个域:Coding、Instruction Following、Math Reasoning、Knowledge Recall、Safety。512 个子集各 100K 样本,proxy model 是 Qwen2.5-0.5B。下游评估分 Development set 和 Unseen set。
主要对比在三个数据规模(100K/400K/800K)和两个模型规模(0.5B/7B)上做。CausalMix 在 800K + 7B 的设定下,CausalMix-S AvgDev 约 62.28(最优),CausalMix-A AvgUns 约 49.09(最优);RegMix 为 60.14 / 48.12,DMO 为 60.35 / 48.98。两个 variant 各有胜负:S 在 Development set 上更好,A 在 Unseen set 上更好。
更能说明因果建模价值的是迁移实验:把 tulu-3-sft-mixture 上训好的因果模型直接用于 AM-Thinking-v1 的长思维链数据(一个完全不同的数据池),在 Qwen3-4B-Base 上训练和评估。CausalMix 在 Math+Code 平均分 66.66,RegMix 61.40,DoReMi 62.00,DMO 63.47。迁移场景下的提升幅度(约 3 到 5 个点)比主实验更大,说明因果模型对数据池变化的鲁棒性确实比全局回归好。
CATE 可解释性
论文把 CausalForestDML 的树结构做了可视化。几个有意思的发现:
- Instruction Following 数据在几乎所有数据状态下都有正的边际回报,是"安全加"的选择
- Knowledge 数据在高难度、高复杂度的数据状态下边际回报为负,反映出事实知识和逻辑推理之间的"技能冲突"
- Math、Coding、Safety 数据在低质量数据上边际回报为负(引入分布噪声),但在中等质量数据上有协同增益

上图是 CausalForestDML 的树结构可视化。每个节点按 covariates(HES、Normalized_Loss、Writing_Style)分裂,叶节点里的数字是各域的 CATE 估计值(正值表示增加该域比例能提升性能,负值表示负迁移)。根节点按 Normalized_Loss 分裂,说明数据难度是决定配比策略的首要条件;IF 在所有叶节点里都是正值(紫色),Knowledge 在高 HES 区域变负(红色),Math/Coding/Safety 在低 Writing_Style 的叶节点里呈负值。
这些 insight 用 RegMix 的 LightGBM 特征重要性是看不出来的,因为 RegMix 不区分数据状态。
局限性
代理规模仍然受限于 512 个 0.5B 模型。covariates 在三个时效果最好,增加更多后性能下降,论文归因于 meta-dataset 太小(只有 512 条),因果估计器在高维时方差太大。
因果识别假设需要配比是在训练和评估之前就确定的,不能根据训练动态调整。这排除了 online data mixing 类方法的适用场景。
$X$ 的定义目前依赖 OpenDataArena 提供的预计算分数,换一个没有这些分数的数据集需要先跑一遍打分。
小结
CausalMix 相比 RegMix 的改进是从"全局回归"升级到"条件因果效应估计":加入数据状态作为 covariate,通过 DML 正交化把混杂效应剥离,用因果森林估计 CATE。收益有两层:一是在同一数据池上性能稍优(1 到 2 个点),二是能直接迁移到不同数据池而无需重新做代理实验(迁移实验上提升约 3~5 个点)。后者对实际工程的价值更大。
对已经熟悉 RegMix 的读者,CausalMix 可以理解为:RegMix 拟合一条 $T \to Y$ 的全局回归线,CausalMix 在这条线的基础上把"数据长什么样"这个条件变量加了进来,通过因果推断的正交化技术保证加入 $X$ 不会引入新的偏差。