SMAT Simulates Merging at Train Time and Adds 1.07-2.16 Points at Under 2% Cost

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 simulates merging during expert training via Scale, Mask, and Perturb, gaining 1.07-2.16 points over the strongest baseline on four backbones at under 2% extra training cost.

What problem this solves

Model merging adds independently fine-tuned expert updates back onto a shared pretrained init, without joint retraining. Standard fine-tuning only cuts each expert's task loss. When those updates overlap or flip signs, merged accuracy drops. TIES, DARE, and DELLA prune or reweight coordinates at merge time, but the expert never saw those operations while it was training: its own task vector can be scaled down, coordinates can be dropped, and other experts' updates get added on top.

Prior merge-aware training (MAT) only covers part of that. SAFT flattens the loss near the expert weights. MergOPT injects noise to stand in for other experts. OrthoReg pushes update matrices toward orthogonality. None of them fully treat scaling and masking of the expert's own update, and they pay for extra forward passes, perturbation machinery, or matrix regularizers. On Llama-3.2-1B, OrthoReg takes 5.69× the wall time of ordinary fine-tuning; ASAM takes 2.87×.

That training cost is high enough that merge-time pruning cannot buy it back.

Method

PolyU, HKUST Guangzhou, CUHK, and InfiX treat Task Arithmetic, TIES, DARE, and DELLA as three operations an expert can simulate on its own. A task vector is the expert weights minus the pretrained weights, which is what merging actually adds.

The objective mixes expert loss with expected loss at the simulated merged point. One sampled state is a stochastic estimate of that expectation; training does not enumerate merge coefficients or expert combinations. Each step evaluates only one of the two losses: a cycle of t=4 does three ordinary steps and one simulated-loss step, so one forward and one backward per update. Mixed-loss steps can be replaced, to first order, by a short run of single-loss steps, which is why the periodic schedule is cheap. The simulated gradient is then pointwise-multiplied by α and the mask before it hits the expert; dropped coordinates get no credit from that loss. Two Triton kernels fuse the three ops and the gradient rescaling. Original and simulated weights live in separate buffers; the model switches to the simulated buffer for the pass and back before the optimizer step. Language runs use (αmin, σ)=(0.2, 2×10⁻³); vision uses (0.1, 10⁻³); p=0.5.

Results

Four backbones, five mergers (weight averaging, Task Arithmetic, TIES, DARE then TA, DELLA), against FT, ASAM, MergOPT, and OrthoReg. Language follows MergOPT's TRACE recipe; vision follows FusionBench's CLIP fine-tune setup.

BackboneSMAT five-merger meanStrongest baselineTime vs FT
Llama-3.2-1B45.77OrthoReg 44.661.019×
Llama-3.1-8B58.49OrthoReg 57.421.014×
CLIP ViT-B/3274.46MergOPT 72.571.010×
CLIP ViT-L/1487.78OrthoReg 85.621.002×

On 1B, SMAT beats OrthoReg by 1.12 points and uses about 82% less training time. On 8B the gap is 1.07. The two CLIP encoders gain 1.89 and 2.16 over their strongest MAT baselines. Training-time overhead versus FT is 0.2%-1.9%; peak GPU memory rises 1.7%-24.3%, from 30.7 GiB to 38.2 GiB on 8B. On 1B, SMAT wins three of five merger columns (TIES, DARE, DELLA); OrthoReg leads weight averaging 43.85 to 42.59, and ASAM leads Task Arithmetic 46.29 to 46.11. On 8B and both CLIP backbones it wins all five columns.

Solo expert scores do not always rise. On 1B they go from 54.43 (FT) to 56.61; on 8B they fall from 62.88 to 61.23. Better merges can come with a weaker standalone expert.

Swapping AdamW for Muon on Llama-1B, SMAT's five-merger mean is 44.15, 0.74 above OrthoReg and 5.90 above FT, still at 1.00× FT time. Ablating Perturb, Scale, and Mask costs 2.10, 1.25, and 1.07 points. As the expert count grows from 2 to 7, the score normalized to FT experts drops from 93.20% to 86.08%, but SMAT leads all three baselines at K=5, 6, and 7. Replacing FT experts with SMAT experts one task at a time lifts that normalized score from 77.55% (all FT) to 86.08% (all SMAT). Loss slices along scaling the expert's own update and adding other tasks' updates show a wider low-loss basin than FT or MergOPT.

Why it matters

For teams already merging task vectors, this changes the training objective and leaves the merger as-is. Cost stays close to ordinary fine-tuning, the same experts score higher on average across five mergers, and the code is public. If the workflow is to fine-tune separately and pick a merger later, this is a training recipe that avoids OrthoReg's 5-11× wall-clock hit.

The lift is a 1-2 point incremental gain. The 8B standalone expert still dips. Time looks cheap; 8B memory grows by about a quarter, which is not free at scale.

Limitations

The paper states that SMAT assumes shared architecture and initialization. Heterogeneous fusion, low-bit, calibration, and experts that arrive over time are left as future work. The appendix also notes that a smaller expected-loss gap versus a real merge distribution does not imply higher task scores.

Perturb replaces real task vectors with isotropic uniform noise. The three noise shapes were compared only on 1B, so it is unclear whether that proxy holds for larger models or strongly coupled tasks. Merge coefficients and sparsity are tuned on a held-out dev split, with sparsity calibrated on FT and then shared across methods, which may assign part of the gain to the training method. The 8B expert-score drop is not analyzed on its own. The bound that justifies swapping mixed-loss steps for a periodic schedule covers one cycle of plain gradient descent; AdamW and Muon adaptive state are outside that argument.

Terms

Source

Related papers

All paper explainers