SMAT는 병합 연산을 학습에 시뮬레이션해 병합 성능을 끌어올린다

SMAT: Simple and Efficient Merge-Aware Training

HF Daily2609.33437

Yanggan Gu, Yuanyi Wang, Zhen Li2026-09-27조회 2

무엇인가

모델 병합은 같은 사전학습 초기값에서 출발한 여러 전문가의 파라미터 업데이트를 합쳐, 공동 재학습 없이 여러 능력을 한 모델에 담는 기법이다. 문제는 표준 파인튜닝이 자기 태스크 손실만 최소화한다는 점이다. 전문가들의 업데이트가 겹치거나 부호가 충돌하면 병합 후 성능이 떨어진다. 이를 학습 단계에서 해결하려는 계열이 병합 인지 학습(MAT)이다. 기존 MAT인 SAFT는 전문가 가중치 근처에서 낮은 손실을, MergOPT는 병합을 흉내 낸 노이즈를, OrthoReg는 업데이트 행렬 내 직교성을 각각 노린다. 저자들은 이들이 병합에서 실제로 일어나는 스케일링과 마스킹, 특히 전문가 자기 자신의 업데이트에 가해지는 변화를 충분히 반영하지 못하고 추가 연산 비용을 만든다고 지적한다.

어떻게 동작하나

논문의 출발점은 관찰이다. Task Arithmetic, TIES-Merging, DARE, DELLA 같은 대표적 병합 기법은 전문가 i의 관점에서 세 연산으로 환원된다. Scale은 자기 업데이트 Δi에 가중치 αi를 곱하고, Mask는 선택된 좌표를 제거하며, Perturb는 다른 전문가들의 업데이트 합 εi를 더한다. 즉 병합 파라미터는 θ0 + αi⊙mi⊙Δi + εi로 쓸 수 있고, 전문가 i 기준 변화량은 (αi⊙mi − 1)⊙Δi + εi가 된다. 여기서 ⊙는 좌표별 곱셈이다. SMAT는 이 세 변화를 독립 학습 중에 무작위로 샘플링해 시뮬레이션한다.

무엇과 다른가

구체적으로 Scale은 αi ~ U[αmin, 1]에서 스칼라 계수 하나를 뽑아 모든 좌표에 공유한다. Mask는 DARE식 무작위 드롭과 재스케일링을 따르며, 기본적으로 트랜스포머 어텐션과 MLP의 선형 가중치에만 적용하고 편향·정규화 파라미터·임베딩은 제외한다. 좌표 k는 확률 pk로 0이 되고 그렇지 않으면 1/(1−pk)로 커지는데, 기댓값이 1이라 스케일된 업데이트의 기댓값은 보존된다. Perturb는 다른 전문가의 체크포인트에 접근할 수 없으니 무작위 파라미터 노이즈로 대체한다. 저자들은 동일 RMS에서 균등·가우시안·라플라스 노이즈를 비교했고, RMS 2×10⁻³에서 균등 분포가 가장 좋았으며 라플라스는 일관된 이점이 없었다고 보고한다. 최종 목적함수는 (1−λ)·전문가 손실 + λ·시뮬레이션된 병합 파라미터에서의 기대 손실이다. 역전파는 ∇θi L(θ̃i) = αi·mi ⊙ ∇θ̃i L(θ̃i) 형태가 되어, 병합에서 배제된 좌표는 시뮬레이션 손실로부터 그래디언트를 받지 않는다.

어떻게 쓰나

효율화가 이 논문의 절반이다. 저자들은 혼합 손실 업데이트를 각 손실에 대한 단일 손실 스텝으로 근사할 수 있음을 명제 1로 제시한다(두 그래디언트가 립시츠이고 노름이 유계일 때 파라미터 차이가 η²LgG·n(n−1) 이하). 이를 근거로 주기적 스케줄링을 쓴다. 한 주기에 전문가 손실 업데이트 t−1회, 시뮬레이션 손실 업데이트 1회로 λ=1/t에 대응하며, 실험에서는 t=4를 쓴다. 각 스텝은 순전파 1회와 역전파 1회만 수행한다. 여기에 두 개의 Triton 커널로 가중치 시뮬레이션과 그래디언트 재스케일링을 융합해 파라미터 크기 텐서의 반복 읽기·쓰기를 줄이고, 시뮬레이션 가중치를 재사용 버퍼에 두었다가 옵티마이저 업데이트 전에 원래 저장소로 되돌리는 파라미터 저장소 전환을 적용한다. 고정된 베이스 가중치는 pinned CPU 메모리에서 프리페치하며, 추가 GPU 버퍼는 학습 가능 가중치 크기 하나만 필요하다.

전제와 한계

실험은 언어와 비전-언어 네 개 백본에서 이뤄졌다. Llama-3.2-1B-Instruct와 Llama-3.1-8B-Instruct를 TRACE의 7개 태스크(C-STANCE, FOMC, MeetingBank, ScienceQA, NumGLUE-cm, NumGLUE-ds, 20Minuten)에서, CLIP ViT-B/32와 ViT-L/14를 Cars, DTD, EuroSAT, GTSRB, MNIST, RESISC45, SUN397, SVHN에서 평가했다. 비교 대상은 표준 파인튜닝(FT), ASAM(SAFT 구현), MergOPT, OrthoReg이고, 병합은 가중치 평균(WA), Task Arithmetic(TA), TIES, DARE+TA, DELLA 다섯 가지다. 5개 병합 평균 점수는 Llama-1B에서 45.77, Llama-8B에서 58.49로 최강 베이스라인 OrthoReg를 각각 1.12점, 1.07점 앞섰다. 1B에서는 5개 병합 열 중 3개, 8B에서는 5개 전부에서 앞선다. 비전에서는 ViT-B/32가 74.46으로 MergOPT(72.57)를 1.89점, ViT-L/14가 87.78로 OrthoReg(85.62)를 2.16점 앞서며 두 백본 모두 5개 열 전부 1위다. 학습 시간 오버헤드는 네 백본에서 0.2~1.9%, 최대 GPU 메모리는 1.7~24.3% 증가했다. Llama-3.2-1B에서는 OrthoReg보다 1.12점 높으면서 학습 시간은 82% 적게 썼다.

추가 분석도 여럿 있다. 옵티마이저를 Muon으로 바꿔도 5개 병합 평균 44.15로 OrthoReg를 0.74점, FT를 5.90점 앞섰고 학습 시간은 FT의 1.00배였다. 전문가 수를 2~7개로 늘린 TA 실험에서 SMAT는 전문가가 적을 때는 MergOPT와 비슷하지만 5·6·7개에서는 세 베이스라인을 모두 앞섰고, 정규화 점수는 2개일 때 93.20%에서 7개일 때 86.08%로 떨어졌다. FT 전문가 7개 중 일부를 교체하는 실험에서는 전부 FT일 때 77.55%가 전부 MergOPT 80.97%, 전부 OrthoReg 83.51%, 전부 SMAT 86.08%로 올랐다. 절제 실험에서 Scale·Mask·Perturb를 각각 제거하면 평균이 1.25점, 1.07점, 2.10점 떨어져 Perturb의 기여가 가장 컸다. 저자들은 병합 방향을 따라 손실을 측정해 SMAT가 더 넓은 저손실 영역을 만든다는 것을 확인하고, 이를 Perturb의 손실 평활화 효과로 해석한다.

실무 관점에서 SMAT는 이미 개별 파인튜닝 후 병합하는 파이프라인을 쓰는 팀에 바로 끼워 넣을 수 있는 형태다. 학습 루프에 시뮬레이션 스텝을 4분의 1 비율로 섞고, αmin·σ·마스크 확률 p 같은 하이퍼파라미터만 정하면 되며, 언어 백본은 (αmin, σ)=(0.2, 2×10⁻³), 비전 백본은 (0.1, 10⁻³), p=0.5를 썼다. 다만 저자들이 명시한 전제를 확인해야 한다. Perturb는 다른 전문가의 체크포인트와 태스크 벡터를 쓸 수 없다는 전제에서 무작위 노이즈로 근사한 것이고, 손실 평활화 해석은 부록의 매끄러움 가정에 의존한다. 목적함수가 실제 병합을 얼마나 잘 대표하는지는 샘플링된 상태의 평균과 2차 모멘트에 달려 있다. 마스크는 편향·정규화 파라미터·임베딩을 제외하고, 태스크별 헤드는 병합하지 않고 분리해 둔다. 또한 8B 언어 모델에서는 전문가 자체 성능이 FT보다 낮아졌는데, 이는 병합 성능 향상이 독립 전문가 성능 향상과 반드시 일치하지는 않는다는 점을 보여준다.