ReLU 유닛별 투영으로 파인튜닝 망각을 줄인다
A Tropical Geometry View of Forgetting: A Per-Unit Projector for Knowledge-Preserving Fine-Tuning
무엇인가
파인튜닝을 하면 모델이 원래 하던 일을 잃는다. 재학습(replay) 없이 이를 막는 대표적 방법인 Adam-NSCL과 GPM은 레이어 입력의 단일 부분공간을 업데이트의 모든 행에서 금지한다. 이 논문은 그 제약이 너무 거칠다고 주장한다. ReLU 유닛은 대부분의 (토큰, 유닛) 쌍에서 닫혀 있어 출력이 정확히 0이고, OPT-1.3b에서는 그 비율이 96%다. 공유 부분공간은 이 닫힌 쌍들까지 전부 제약하지만, 닫힌 쌍은 가중치가 움직여도 그 자리에서 계속 0을 낸다.
어떻게 동작하나
이론적 뼈대는 ReLU 레이어의 두 쌍대 그림이다. 데이터 공간에서 유닛의 벽은 tropical 초곡면이고, 게이트 패턴의 세포들은 zonotope의 상부 꼭짓점과 쌍대다. 가중치 공간에서는 옛 토큰 하나하나가 초평면이 되고, 토큰들이 모든 토큰을 자기 편에 남겨두는 가중치들의 폐포인 다면체 P_X를 잘라낸다. Theorem 5는 두 그림을 잇는 정확한 항등식이다. 어떤 가중치 변화에 대해서도 레이어 출력의 제곱 변화는 in-cell, open에서 closed로, closed에서 open으로 가는 세 항으로 나뉘고 나머지 항이 없다. 앞의 두 항은 각 유닛이 발화하는 토큰, 즉 그 유닛의 open 토큰 위에서만 살아 있다.
무엇과 다른가
여기서 나오는 것이 gate-aware per-unit projector다. 유닛 i마다 open 토큰 중 post-activation이 큰 상위 C개를 골라 행렬 A_i를 만들고, 옵티마이저 스텝마다 누적 변위 u_i를 u_i − A_i^T(A_iA_i^T+ρI)^{-1}A_i u_i로 치환한다. ρ는 A_iA_i^T 평균 대각의 1e-4배이고, ρ=0 극한에서 그 행은 rank A_i ≤ k_i개의 방향을 잃는다. Theorem 9는 보호 비용을 계산한다. 유닛 i를 정확히 보호하려면 그 유닛 자신의 open 토큰 랭크 r_i만큼의 방향이 필요하지만, 모든 행이 공유하는 하나의 부분공간은 매 행에서 open 토큰 합집합의 랭크 이상을 지불해야 한다. Proposition 10은 예산이 제한됐을 때 최적의 per-unit 제약이 그 유닛 자신의 open-token Gram 행렬의 top-q 고유공간임을 보인다. Proposition 12는 그래디언트나 Fisher 같은 1차 기준이 닫힌 쌍에서 0이 되기 때문에 이 open/closed 분할을 자동으로 물려받는다는 것을 보인다.
어떻게 쓰나
실험은 OPT-125m, OPT-1.3b, OPT-6.7b(ReLU, θ=0, 6.7b는 마지막 네 블록의 fc1)에서 한다. 보존할 지식은 MiniPile로 대표되는 사전학습 분포이고 held-out loss와 5개 제로샷 태스크로 채점한다. 2048토큰 support로 projector를 만들고 겹치지 않는 1024토큰 probe로 망각을 잰다. 새 과제는 WikiText-103 또는 CodeParrot의 Python 코드이고, fc1을 65,536토큰에 AdamW lr 1e-4, cosine, 128스텝, 배치 4로 학습한다. OPT-1.3b에서 per-unit projector는 행당 9~60개 제약 방향의 여섯 예산 모두에서 Adam-NSCL보다 망각이 적고, 27개 seed-pair 중 24개에서 이긴다(부호 뒤집기 순열검정 p=1.5×10⁻⁶, 랜덤효과 차이 −0.018 nats, 95% CI [−0.028, −0.009]). Adam-NSCL이 더 배우는 양은 최대 0.003 nats다. 격차는 행당 9방향에서 1.1배, 60방향에서 4.3배로 벌어진다. 공유 부분공간은 행당 15~47방향 구간에서 0.046~0.050 nats로 정체하는데 per-unit은 0.010까지 계속 내려간다. GPM의 에너지 기준이 고르는 평균 259방향과 맞붙였을 때 cap 256(행당 47.3방향)은 2.3배 덜 잊고(3/3, p=0.009) 0.052 nats 더 배운다. cap 256은 OPT-1.3b의 비제약 망각을 95.9%, OPT-125m에서는 97.7% 제거한다. OPT-6.7b에서는 cap 128이 행당 68.7방향을 제약해 Adam-NSCL k_sh=69와 같은 망각(0.0192 대 0.0184, TOST 동등)을 내면서 0.011 nats 더 배운다(3/3, p=1.5×10⁻⁴). 둘 다 비제약 망각 0.126의 85%를 제거한다.
전제와 한계
무엇이 기준인지 분리하는 실험도 있다. 유닛마다 k_i를 23.5로 고정하고 어떤 토큰으로 채우는지만 바꾸면, 그 유닛 자신의 open 토큰이 랭크를 맞춘 랜덤 토큰, |z−θ|로 정렬한 부호맹 토큰, 가장 닫힌 토큰을 18/18 seed-pair에서 이긴다(p=3.8×10⁻⁶). 다만 open 집합 안에서 무작위로 뽑은 토큰도 가장 활성인 토큰과 비슷하게 보호한다. 두 집합은 쌍의 30%에서 다르고 활성 에너지의 35.2% 대 59.6%를 덮는다. 즉 기준은 순위가 아니라 open/closed 분할 자체다. 1차 기준도 |∂L/∂z|가 고르는 쌍의 99.6%가 open이라 같은 결론에 도달한다. 반대로 유닛을 통째로 동결하는 것은 훨씬 약하다. 같은 23.5방향에서 per-unit이 0.025를 잊을 때 행 동결은 0.210을 잊고(6/6, p<10⁻⁴) 0.039 nats 덜 배운다. 출판된 기법들과 비교하면 EWC는 λ=4에서 4×10⁴까지 지배당하고 λ=4×10⁵에서 per-unit frontier에 닿으며, MIGU는 기본 마스크 비율 0.7에서 cap 128보다 2배 더 잊고 0.051 nats 덜 배운다(p=0.0005). LoRA는 망각에서 5배 안에 들어오지 못하고, LoRA-Null은 LoRA의 망각을 2.3~3배 줄이지만 여전히 cap 128보다 3.6~5.1배 더 잊는다. WikiText-103 비제약 파인튜닝이 OPT-1.3b의 5개 제로샷 평균 정확도를 2.94점 떨어뜨리는 데 비해 cap 128은 1.02점만 떨어뜨린다(3/3, p=0.013).
두 번째 응용은 프루닝 후 수리다. Wanda 같은 마스크로 가중치를 0으로 만든 뒤 남은 가중치를 고쳐 레이어가 옛 출력을 재현하게 하는데, 논문은 닫힌 캘리브레이션 토큰이 임계값을 넘어가는 양 Esc(ŵ)를 정의한다. Lemma 13은 밀집 가중치에서 출력 오차의 모든 미분 기반 국소 모델이 closed에서 open으로 가는 항에 눈이 멀다는 것을 보인다. 그래서 gate-weighted 목적식의 최소화점이 다면체 P_X를 벗어날 수 있고(Theorem 14), 그 닫힌 형태 해는 OPT-1.3b, sparsity 0.7에서 수리하지 않는 것보다 1.94 nats 나쁘다. Theorem 15는 한쪽 페널티 λ Esc를 더하면 탈출이 ||ζ||²/λ + E_0 이하로 묶인다고 말한다.
개발자 관점에서 이 논문이 말하는 바는 분명하다. 도메인 적응을 위해 LoRA나 전체 파인튜닝을 할 때 옛 능력이 얼마나 깎이는지가 문제라면, 제약을 레이어 전체 공유 부분공간이 아니라 유닛별 open 토큰에 걸어야 한다. 구현도 무겁지 않다. base 모델로 옛 토큰 support를 한 번 forward해 유닛별 post-activation을 기록하고, 유닛마다 상위 C개 토큰의 입력 행렬 A_i를 만들어 둔 다음, 옵티마이저 스텝마다 그 유닛 행의 누적 변위를 A_i의 ridge 영공간으로 사영하면 된다. 그래디언트가 필요 없고 forward 한 번이면 된다. 다만 이득은 open 토큰이 충분히 덮일 때 나온다. OPT-125m처럼 유닛이 많은 토큰에서 발화하면(2048개 중 평균 111개, 1.3b는 72개) 같은 cap으로는 공유 부분공간이 더 나을 수 있고, 행당 17방향 근처에서야 역전된다. cap이 각 유닛 open 토큰 span의 대부분을 덮는지부터 확인해야 한다.
한계도 분명히 밝혀져 있다. 논문의 정리들은 한 레이어를 그 레이어의 base 입력에서 본 것만 말한다. 실제로는 모든 fc1 레이어가 동시에 학습되므로 end-to-end 수치에는 각 레이어 입력의 드리프트가 섞인다. ρ>0일 때 방향은 제거되는 게 아니라 축소된다. 이론은 정확 보호에 도달하기 전 구간에서 per-unit과 Adam-NSCL의 우열을 정하지 못하며, P1은 그 구간에 대한 예측일 뿐이다. 항등식의 쌍별 케이스 분할은 27개 pruned OPT 셀에서 상대 오차 5.8×10⁻⁵, 두 7B ReLU 계열 모델의 18개 셀에서 9.0×10⁻⁵까지 성립한다. 프루닝 수리 쪽에서 Q_M이 비어 있지 않음을 보증하는 bias witness는 측정된 wall 유닛의 81~96%에서만 확인된다.