순열 등변 플로우 매칭으로 정렬 없이 신경망 가중치를 생성하는 방법
Permutation-Equivariant Flow Matching for Alignment-Free Neural Weight Generation
무엇인가
학습된 신경망은 고차원 파라미터 벡터 하나로 표현된다. 이 벡터들의 분포를 학습하면 최적화 대신 샘플링으로 새 모델을 얻을 수 있다는 것이 weight-space 생성 모델링의 야심이다. 문제는 순열 대칭이다. 은닉 뉴런의 순서를 바꾸면 함수는 완전히 같지만 가중치 벡터는 전혀 다른 위치로 이동한다. Hugging Face 같은 저장소에는 공통 순서나 공유 초기화 없이 독립적으로 학습된 네트워크가 쌓여 있으므로, 생성 모델 입장에서는 같은 함수가 여러 개의 먼 점으로 흩어져 보인다. 기존 연구는 공통 베이스 모델에서 파생된 체크포인트를 모으거나, Git Re-Basin 같은 근사 정렬로 사후에 뉴런을 맞추는 방식으로 이 문제를 우회했다. 전자는 아키텍처가 호환되어야 하고, 후자는 정렬 비용과 근사 오차를 감수해야 하며, 둘 다 이미 존재하는 독립 학습 컬렉션에는 그대로 적용하기 어렵다.
어떻게 동작하나
이 논문은 정렬을 아예 하지 않는다. 신경망을 파라미터 그래프로 표현한다. 노드는 뉴런이나 채널, 엣지는 그 사이의 가중치를 담고, 위치 인코딩이 레이어 소속 같은 구조적 역할을 알려준다. 이 그래프 위에서 Flow Matching을 수행하되, 노이즈에서 학습된 가중치 분포로 이동시키는 속도장 v_θ(G, w_t, t, c)를 순열 등변 Graph Meta Network(GMN)로 parameterize한다. 시간 t와 태스크·도메인 조건 c는 AdaLN으로 각 메시지 패싱 블록에 주입되고, 노드·엣지·전역 특징이 순열 불변 aggregator를 거쳐 갱신되며, 공유 readout MLP가 각 엣지의 속도(즉 각 가중치의 변화량)를 예측한다. 학습 목표는 예측 속도와 실제 (w_1 - w_0)의 L2 오차다. 저자들은 이 설계가 임의의 선택이 아니라는 근거도 제시한다. 초기화 분포가 순열 불변이고 학습 절차가 등변이면 결과 가중치 분포도 불변이며, 사전분포와 목표 분포가 모두 불변일 때 FM 최적 속도장은 거의 확실히 등변이다. 게다가 등변 속도장에 불변 사전분포와 독립 커플링을 쓰면 각 학습 네트워크에 서로 다른 순열을 적용해도 기대 목적함수가 변하지 않는다. 즉 정렬 전처리는 도움이나 해가 될 수 없으므로 생략하는 편이 낫다는 것이다.
무엇과 다른가
평가도 새롭다. 정확도만 높으면 학습 체크포인트를 그대로 베낸 것일 수 있고, 유사도만 낮으면 기능하지 않는 출력일 수 있다. 그래서 테스트 정확도(또는 ROC AUC)인 task performance, 최대 error-IoU인 functional similarity, Git Re-Basin으로 정렬한 뒤의 최대 코사인 유사도인 WCS를 한 네트워크당 3차원 점수 벡터로 묶고, 생성 네트워크 집합과 학습 컬렉션 부분집합 사이의 Wasserstein-2 거리를 재서 Joint Wasserstein Similarity(JWS)로 환산한다. 이 결합 지표가 논문의 주 정량 지표다.
어떻게 쓰나
무조건부 생성 실험은 두 개의 독립 학습 컬렉션에서 이뤄졌다. 두 은닉층 MNIST MLP 1만 개와 3개 합성곱층 + 선형 분류기 CIFAR-10 CNN 5만 개이며, 모두 초기화 시드가 다르다. 비교 대상은 DWF, P-diff, SANE, 그리고 MLP·DiT 속도망이다. 결과는 JWS 0.91/0.92로, 베이스라인 최고치 0.59/0.64를 크게 앞선다. MLP, P-diff, 정렬하지 않은 DWF는 사실상 chance 수준이었고, 정렬은 DWF를 도왔지만 정확도는 참조 분포에 못 미쳤다. SANE의 가우시안 모드는 성능이 나빴고, KDE 샘플링과 단순 섭동은 과도한 가중치 유사도에도 정확도를 유지했다. 저자들은 이를 정확도만 보면 안 되는 이유로 든다. 같은 MNIST GMN을 Git Re-Basin으로 정렬한 컬렉션에서 다시 학습시켜도 정확도와 유사도 통계가 거의 동일했다는 점도 보고된다.
전제와 한계
스케일과 조건부 생성도 검증했다. mfeat-karhunen에서 학습 컬렉션을 100개에서 5만 개까지 늘리며 2M 파라미터 GMN과 2M·10M DWF/DiT를 비교했는데, 컬렉션이 작을 때 베이스라인의 높은 정확도는 Max-IoU와 WCS가 1에 가까워지는 암기 현상과 함께 나타났고, 2M GMN이 전 구간에서 가장 높은 JWS를 기록했다. 조건부 설정에서는 OpenML의 테이블 데이터 20종, 16개 아키텍처(입력 차원·출력 클래스·깊이·은닉 폭이 제각각)에 걸친 10만 개 독립 학습 MLP를 단일 조건부 GNN 플로우로 학습했다. 생성 네트워크의 평균 테스트 정확도는 82.17%로 컬렉션의 83.15%에 근접했고, 학습 때 본 적 없는 은닉 폭 0.5배·1.5배 구성으로의 제로샷 전이에서도 각각 78.96%, 82.35%를 얻었다. letter 태스크가 원본 대비 격차가 가장 컸고, 폭 축소는 letter와 mfeat-karhunen에 가장 큰 영향을 줬다.
도메인 시프트 실험은 Folktables/ACSIncome의 캘리포니아(CA)와 푸에르토리코(PR) 데이터로 했다. 연소득 5만 달러 초과를 예측하는 분류기를 도메인당 5천 개씩 독립 학습시키고, 조건부 플로우는 c=0(CA)과 c=1(PR) 두 끝점에서만 학습한 뒤 학습에 없던 c∈(0,1)에서 샘플링했다. 중간 조건의 생성 네트워크는 CA-PR 성능 트레이드오프를 따라 이동하며 두 네트워크 로짓 앙상블에 필적하는 cross-domain AUC를 냈다. 정렬하지 않은 가중치 평균은 성능을 떨어뜨렸고 사후 정렬은 혼합 성능을 개선했다. JWS는 양 끝점의 0.969/0.976에서 c=0.5에서 0.700/0.824로 낮아졌는데, 이는 도메인 참조 컬렉션에서 멀어졌다는 뜻으로 해석된다.
개발자 입장에서 이 논문이 흥미로운 지점은 전처리 파이프라인을 하나 없앤다는 것이다. 뉴런 정렬은 아키텍처가 맞아야 하고 비용도 크며, 체크포인트 궤적을 직접 만들어야 하는 기존 방식은 이미 공개된 모델 저장소에는 쓸 수 없다. 순열 등변 구조를 생성기 자체에 심으면 독립 학습 모델을 그대로 학습 데이터로 삼을 수 있고, 하나의 조건부 모델이 서로 다른 아키텍처와 본 적 없는 은닉 폭까지 처리한다. 실무에서 응용을 검토한다면 세 가지를 확인해야 한다. 첫째, 지금 결과는 소형 네트워크 규모이므로 실제 타깃 모델 크기에서 메모리·연산이 감당 가능한지. 둘째, 조건화가 데이터셋 특징이 아니라 태스크 식별자에 의존하므로 새 태스크로의 일반화는 보장되지 않는다는 점. 셋째, 정확도 하나만 보지 말고 JWS처럼 성능과 유사도를 함께 보는 지표로 암기 여부를 점검하라는 것이다.
저자들이 명시한 한계도 분명하다. 실험은 의도적으로 작은 네트워크에 집중했고, 이는 확장성 이전에 소규모에서 암기 없이 생성이 되어야 한다는 판단에서다. dense 그래프 표현은 상당한 메모리와 연산을 요구해 대형 모델 확장은 향후 과제로 남으며, 커스텀 GPU 커널이나 최적화된 희소 연산으로 완화할 수 있다고 본다. 그래프 표현이 Transformer로 확장되기는 하지만 파운데이션 모델 가중치 생성은 여전히 열린 문제다. 또한 조건화가 데이터셋 특징이 아닌 태스크 식별자를 쓰고 있어 본 적 없는 태스크에서의 생성은 평가하지 않았으며, Set Transformer나 DeepSets로 데이터셋 표현에 조건화하는 방향이 남아 있다. JWS 외에 weight-space 생성을 더 폭넓게 평가할 보완 지표도 필요하다고 밝힌다.