SMAT: Simple and Efficient Merge-Aware Training
Yanggan Gu, Yuanyi Wang, Zhen Li, Shuo Cai, Yuhang Liu, Junzhuo Li, Zihao Wang, Hongxia Yang
cs.LG, cs.CL
2026-09-27
SMAT在专家训练中用缩放、掩码和噪声模拟合并,四类骨干上五类合并均值比最强基线高1.07-2.16分,相对微调开销不到2%。
模型合并把多个独立微调专家的参数更新叠回同一个预训练初始化,省掉联合重训。专家训练只盯自己的任务损失,叠在一起时更新会重叠、符号会打架,合并后单任务分数往下掉。TIES、DARE、DELLA 在合并阶段按幅度、符号或随机丢掉一部分坐标,能减轻冲突,但专家在训练时并不知道自己的更新之后会被缩小、被挖掉、再被别人的更新顶开。
已有的 merge-aware training(MAT,训练阶段就为合并做准备)覆盖也不全。SAFT 在专家权重附近压损失曲率;MergOPT 加噪声模拟别人带来的偏移;OrthoReg 让更新矩阵更正交。它们很少显式处理「自己的 task vector 被缩放和掩码」,还要额外前向、扰动处理或矩阵正则。Llama-3.2-1B 上 OrthoReg 的训练时间是标准微调的 5.69 倍,ASAM 是 2.87 倍。
训练侧为合并付的税太高,合并阶段再怎么挑坐标也补不回来。
港理工、港科广、中大和 InfiX 把常见合并收成三个从单个专家能模拟的操作。task vector 是专家参数减去预训练参数,合并时真正被加减的就是它。
训练目标是 (1-λ)×专家损失 + λ×模拟合并点上的期望损失。一条模拟状态就是对这个期望的一次随机估计,并不枚举合并系数或专家组合。每一步只算其中一项:周期 t=4,三步普通微调、一步走模拟损失,前向和反向各一次。混合损失的若干步可以用分开的专家损失步和模拟损失步近似,误差是学习率的二阶,所以把两类损失按周期拆开算,说得通。模拟损失的梯度回传到专家时再点乘 α 和 mask,被丢掉的坐标从这条损失拿不到更新。fused kernel 把三操作和梯度缩放并成两次 Triton 调用;原权重和模拟权重分两块存储,前向反向用模拟块,优化器更新前切回原块,少一次整模拷贝。语言骨干 (αmin, σ)=(0.2, 2×10⁻³),视觉 (0.1, 10⁻³),mask 概率 p=0.5。
四个骨干、五种合并器(权重平均、Task Arithmetic、TIES、DARE 再接 TA、DELLA),对照 FT、ASAM、MergOPT、OrthoReg。语言跟 MergOPT 的 TRACE 设置,视觉跟 FusionBench 的 CLIP 微调设置。
| 骨干 | SMAT五方法均值 | 最强基线 | 相对微调时间 |
| Llama-3.2-1B | 45.77 | OrthoReg 44.66 | 1.019× |
| Llama-3.1-8B | 58.49 | OrthoReg 57.42 | 1.014× |
| CLIP ViT-B/32 | 74.46 | MergOPT 72.57 | 1.010× |
| CLIP ViT-L/14 | 87.78 | OrthoReg 85.62 | 1.002× |
1B 上比 OrthoReg 高 1.12 分,训练时间少约 82%;8B 高 1.07 分。视觉两条骨干分别高 1.89 和 2.16 分。相对微调,时间开销 0.2%-1.9%,峰值显存多 1.7%-24.3%,8B 从 30.7 GiB 涨到 38.2 GiB。1B 上五种合并里 SMAT 拿下 TIES、DARE、DELLA 三列;权重平均是 OrthoReg 的 43.85 对 SMAT 的 42.59,Task Arithmetic 是 ASAM 的 46.29 对 SMAT 的 46.11。8B 和两条 CLIP 骨干则是五列全赢。
独立专家分数并不总涨。1B 从 FT 的 54.43 到 56.61;8B 从 62.88 掉到 61.23。合并更好,单任务专家可以更差。
Llama-1B 换成 Muon 之后,五方法均值 44.15,比 OrthoReg 高 0.74、比 FT 高 5.90,时间仍是 1.00×。消融里去掉 Perturb 掉 2.10 分,去掉 Scale 掉 1.25,去掉 Mask 掉 1.07。专家数从 2 增到 7,相对 FT 专家的归一化分从 93.20% 降到 86.08%,但在 5、6、7 个专家时都压过三条基线。七个 FT 专家按任务逐步换成 SMAT 专家,归一化分从全 FT 的 77.55% 升到全换的 86.08%。损失切片沿「缩放自己的更新」和「加上别人的更新」两个方向,SMAT 的低损失区域比 FT 和 MergOPT 更宽。
已经在用 task vector 合多专家的团队,这篇改的是训练目标,合并算法可以不动。代价接近普通微调,同一套专家在五种合并器上均值都更高,代码已开源。对「先各自微调、事后再决定怎么合」的流程,这是能换上去的训练配方,不必上 OrthoReg 那种 5-11 倍墙钟的正则。
幅度是 1-2 分的渐进改进。8B 单任务专家还略掉分。时间账好看,8B 显存多约四分之一,大模型上未必免费。
论文写明 SMAT 假定专家同架构、同初始化。异构融合、低比特、校准、专家按顺序到达,都列为未来工作。附录也写了:模拟分布和真实合并分布的损失差,不保证任务分数更高。
Perturb 用各向同性均匀噪声顶替真实 task vector,只在 1B 上比过三种噪声形状,更大模型和强相关任务够不够用,没有直接证据。合并系数和稀疏度在 dev 上搜,稀疏设置按 FT 标定后在方法间共用,可能把一部分收益记到训练方法头上。8B 专家分下降没有单独分析。周期替换混合损失的误差界只覆盖普通梯度下降的一个周期,AdamW 和 Muon 的自适应动量没有对应证明。