DiT 잔차 연결을 단계 대응 검색 경로로 바꾸면 학습이 빨라진다
Structured Residual Connectivity Matters for Diffusion Transformers
무엇인가
Diffusion Transformer(DiT)는 고해상도 이미지 생성의 확장 가능한 백본으로 자리 잡았지만, U-Net 기반 확산 모델이 손으로 설계한 스킵 연결로 다중 스케일 특징을 보존하는 것과 달리 모든 이전 레이어를 하나의 덧셈 상태로 합치는 균일한 잔차 스트림을 쓴다. 이 논문은 이 차이를 문제로 삼는다. 균일한 누적은 개별 레이어의 기여를 점진적으로 희석하고, 역전파 경로에 고정된 토폴로지를 강제해 디노이징 단계별로 그래디언트 흐름을 적응적으로 최적화할 여지를 없앤다는 것이다. 저자들이 던지는 질문은 명확하다. 트랜스포머의 확장성을 유지하면서 생성 모델링에 효과적인 구조적 연결 패턴을 결합할 수 있는가.
어떻게 동작하나
먼저 저자들은 DiT 내부 표현을 체계적으로 분석해 두 가지 관찰을 얻는다. 하나는 인코더 측 의존성으로, 언어 모델에서 흔한 지역적 지배 패턴과 달리 DiT는 네트워크 전체 깊이에 걸쳐 초기 레이어 표현에 높은 중요도를 부여한다. 다른 하나는 자발적 대칭 편향으로, 레이어 간 라우팅 자유도를 주면 모델이 거울 대칭 쌍을 선호한다. 이는 U-Net의 대칭 스킵이 단순한 휴리스틱이 아니라 모델이 스스로 찾아내는 잠재적 구조적 필요성일 수 있음을 시사한다. CKA 표현 유사도 시각화에서도 베이스라인 DiT-XL/2는 깊이에 따라 표현이 강하게 균질화되는 반면, 제안 방법은 과도한 유사도를 낮추고 대각 반대 방향 구조를 더 뚜렷하게 만든다.
무엇과 다른가
제안 방법의 골자는 잔차 연결을 수동적 덧셈에서 능동적 검색 메커니즘으로 바꾸는 것이다. DiT를 깊이 방향으로 인코더 단계와 디코더 단계로 균등 분할하되 원래의 단일 스케일 토큰 해상도는 유지한다. 패치 임베딩 출력 x0을 포함한 인코더 측 중간 표현들을 미분 가능한 잔차 소스 집합으로 수집하는데, 이 소스들은 계산 그래프에서 분리(detach)되지 않아 디코더 측 라우팅의 그래디언트가 인코더 표현으로 직접 흘러간다. 각 DiT 블록에서는 후보 집합에 대해 RMSNorm과 학습된 라우팅 벡터 w로 로짓을 계산하고 소프트맥스로 가중치를 구한 뒤 가중합으로 잔차 소스를 재구성한다. 이 연산자를 셀프 어텐션과 MLP 서브레이어 앞에 모두 삽입해, 잔차 소스가 항등 스트림에 국한되지 않고 선택된 이전 표현의 특징을 끌어올 수 있게 한다.
어떻게 쓰나
소스 선택 전략은 네 가지로 비교된다. 모든 인코더 레이어에 접근하는 전체 라우팅, 첫 번째 또는 마지막 표현만 쓰는 first/last 라우팅, 그리고 각 디코더 레이어가 대응하는 인코더 단계의 표현 하나만 가져오는 미러 라우팅이다. 미러 라우팅은 거울 인덱스 m(j)=K−j+1로 소스를 지정해, 밀집 라우팅의 중복과 모호성을 줄이면서 단계 대응 교차 깊이 상호작용을 보존한다. 실험에서 미러 라우팅이 FID와 IS 모두 가장 좋았고, first/last는 단계 대응이 없어 뒤처졌다.
전제와 한계
실험은 ImageNet 256×256 클래스 조건부 생성, 잠재 공간 표준 프로토콜에서 이뤄졌다. 400K 반복 기준으로 DiT-S/2의 FID를 69.81에서 62.48로, DiT-B/2를 44.21에서 36.33으로, DiT-XL/2를 18.85에서 15.07로 낮췄다. 상대 감소폭은 10.5%, 17.8%, 20.1%로 모델이 커질수록 증가하며, 추가 파라미터는 0.1% 미만이다. DiT-XL/2를 더 오래 학습시키면 800K 반복의 11.40 FID가 베이스라인의 1300K 반복 11.76 FID를 앞선다. 동일 FID 도달에 필요한 학습 스텝은 최대 1.73배 적었다.
이미 강한 사전학습 모델에서도 효과가 나타난다. SiT-XL/2 + REPA는 1M 반복에서 6.4 FID를 찍은 뒤 3M 반복을 더 써서 5.9까지밖에 개선되지 않아 포화 양상을 보였는데, 제안 라우팅을 붙여 0.28M 반복만 파인튜닝하면 4.94 FID, 0.35M 반복에서는 4.34 FID가 된다. 총 1.56 FID 개선이다. 가이던스 구간을 쓰면 0.28M 모델이 1.39 FID에 도달해 REPA의 1.42를 앞선다. 효율 측면에서 DiT-S/2 기준 추가 파라미터는 0.019M, FLOPs 증가는 0.3%에 불과하다. 같은 파라미터 예산의 밀집 Attention Residuals와 FID가 비슷하면서(62.48 대 62.87) 디코더 측 소스 스택 활성화를 약 81% 줄인다(1.51GB에서 289MB).
절제 실험은 융합 방식과 소스 풀의 효과를 분리한다. U-ViT 스타일 정적 미러 스킵은 파라미터가 약 77배 많은데도 400K에서 67.90 FID로 베이스라인 69.81에 근접했지만, 같은 미러 쌍을 콘텐츠 의존 라우팅으로 융합하면 64.83까지 내려가고 인코더 내부 라우팅을 더하면 62.48이 된다. 소스 풀을 인코더 측으로 제한한 Fused-All은 63.48로 밀집 AttnRes의 62.87과 비슷해, 디코더 레이어 출력을 소스로 재사용하는 것이 이미지 생성에는 불필요함을 보여준다. 그래디언트 분석에서는 초기-후기 미러 쌍의 코사인 유사도가 크게 올라가고 중간 레이어는 낮아져, 단순히 모든 레이어를 비슷하게 만드는 것이 아니라 미러 쌍 패턴으로 그래디언트 흐름을 재조직한다.
개발자 관점에서 이 방법은 기존 DiT 학습 파이프라인에 소스 수집과 경량 라우팅 연산자만 추가하는 형태라 도입 비용이 낮고, 특히 REPA처럼 이미 포화된 모델의 후속 파인튜닝에서 여지를 만들어낸다는 점이 실용적이다. 다만 저자들이 밝힌 전제를 확인해야 한다. 실험은 ImageNet 256×256 클래스 조건부 생성에 한정됐고, 인코더/디코더 분할과 미러 대응이라는 구조적 가정 위에서 최적 전략이 선택됐다. 또한 미러 쌍 그래디언트 정렬은 미러 연결 자체가 부분적으로 유발하는 효과이므로, 저자들도 이를 인과적 증명이 아니라 라우팅이 만들어내는 최적화 행동의 특성으로 제시한다. 논문 본문에는 별도의 한계 절이 없고, 결론에서 확장 가능한 생성 모델의 구조적 정보·그래디언트 라우팅에 대한 추가 탐구를 후속 과제로 남긴다.