월드모델이 넓은 잠재벡터에서 예측 정보를 앞쪽 좌표에 몰아넣는 방법
Adaptive Latent Capacity for World Models
무엇인가
JEPA(결합 임베딩 예측 아키텍처) 계열 월드모델은 관측을 재구성하지 않고 미래 임베딩을 예측한다. 문제는 예측에 필요한 정보가 잠재벡터 좌표 전체에 흩어진다는 점이다. 저자들은 임베딩 차원을 키우면 플래닝 성능이 오히려 떨어지는 현상을 관찰했고, 차원을 줄이면 표현력이 제한되고 별도 튜닝이 필요하다고 지적한다. 그래서 넓은 임베딩은 그대로 유지하면서 예측에 유용한 정보를 앞쪽 짧은 프리픽스(첫 몇 개 좌표)에 몰아넣어, 하나의 학습된 모델이 여러 잠재 상태 예산을 지원하도록 만드는 것을 목표로 삼는다.
어떻게 동작하나
핵심 구성은 용량 네트워크다. 허용 가능한 용량 집합 K={k1<…<kC}(kC=d)를 정하고, 마스크 M_k로 앞 k개 좌표만 남긴 z_t(k)=M_k·s_t를 만든다. 용량 네트워크 q_ψ(k|w)는 관측 시퀀스 w를 조건으로 프리픽스 길이에 대한 분포를 출력하며, 학습 중에는 여기서 k를 샘플링해 예측기에 넣는다. 예측기는 마스킹된 현재 임베딩과 행동 a_t를 받아 다음 스텝의 전체 인코더 임베딩 s_{t+1}을 맞히도록 학습된다(손실 L_pred). 목표를 전체 임베딩으로 두기 때문에, 짧은 프리픽스만 주어져도 넓은 표현을 예측하는 데 도움이 되는 정보를 앞쪽 좌표가 보존하게 된다. 용량 분포는 역순위 다항 가중 (C-c+1)^α로 만든 사전 π0로 초기화하고, straight-through Gumbel-Softmax로 미분한다. α>0이면 작은 용량 쪽에, α=0이면 균등, α<0이면 큰 용량 쪽에 무게가 실린다.
무엇과 다른가
정규화 항 MixSIGReg도 새로 도입한다. SIGReg나 VICReg 같은 기존 붕괴 방지 목적함수는 모든 좌표에 균일한 분산을 요구하는데, 뒤쪽 좌표가 자주 꺼지고 마스킹 시 0이 되어야 하는 구조와 맞지 않는다. MixSIGReg는 마스킹된 임베딩을, 활성 프리픽스 구간은 표준 가우시안이고 나머지 좌표는 0인 델타인 혼합 참조분포에 맞춘다. 구현은 임베딩을 단위 방향으로 투영한 뒤 그 혼합분포에 맞춘 Epps-Pulley 특성함수 통계량을 가우시안 주파수 윈도우로 적분해 최소화하는 방식이다. 최종 목적식은 L = L_pred + λ·L_MixSIG다.
어떻게 쓰나
이론적 근거도 제시한다. MixSIGReg 참조분포에서 좌표 j의 분산은 생존확률 ρ_π^j = Pr(K_π ≥ j)와 같고 이 값은 좌표가 뒤로 갈수록 단조 감소한다. 또 프리픽스가 길어질수록 예측 오차는 줄어들 수밖에 없으므로(중첩 정보), 기대 예측 오차는 E[R_K] = R_0 − Σ ρ_π^{k_c} Δ_c로 쓸 수 있다. 각 블록의 예측 이득 Δ_c가 그 블록의 생존확률로 가중되는 구조이므로, 이득이 큰 정보를 앞에 두면 기대 오차가 최소화된다. 다만 저자들은 이 결과가 순서를 유도하는 동기일 뿐, 공동 학습이 실제로 그 순서를 달성한다는 보장은 아니라고 명시한다.
전제와 한계
배포 시에는 MPC를 쓴다. 에피소드의 시작 관측과 목표 관측으로 k_r = argmax q_ψ(k|w)를 골라 에피소드 내내 고정하고, 재귀 롤아웃과 목표 비용 모두 같은 마스크를 쓴다. 비용은 활성 차원만 비교하는 (1/k_r)·||ẑ_H − M_kr·s_g||²이고, CEM으로 후보 행동열을 최적화한다.
통제된 장난감 실험에서는 상태가 4차원(위치·속도 2쌍)인 감쇠 진동자 시스템을 썼다. 관측은 10차원이지만 동역학 요인은 4개뿐이다. 8차원 ALeWM(용량 집합 {1,…,8})의 첫 4개 좌표는 상태 복원 R²=0.983을 달성했는데, 이는 차원을 4로 맞춘 LeWM의 0.980과 비슷하고, 넓힌 LeWM(d=6에서 0.930, d=8에서 0.904)보다 높다. ALeWM은 8차원을 유지하면서도 앞 4개 좌표에 선형 디코딩 가능한 상태 정보를 집중시켰다.
본 실험은 TwoRoom, PushT, Reacher, OGBench-Cube 네 개 목표조건부 시각 제어 벤치마크에서 이뤄졌다. ViT-Tiny(d_max 96·192)와 ViT-Small(d_max 384)을 쓰고, d_max=192일 때 용량 집합은 {8,16,32,64,96,128,160,192}다. 에피소드 단위로 학습·검증·테스트를 분리했고(검증 100, 테스트 200 에피소드), 시드 3개 평균을 보고한다. 결과적으로 ALeWM은 네 과제 모두에서 최고 평균 성공률을 기록했다. 전체 폭 LeWM 대비 성공률이 약 1.5~17%포인트 높으면서 평균 플래닝 용량은 약 67~98% 줄었고, 차원을 튜닝한 LeWM과 비교해도 약 0.3~7.2%포인트 앞서며 12개 비교 중 11개에서 더 적은 플래닝 차원을 사용했다. 가장 큰 개선은 OGBench-Cube에서 나왔다.
절제 실험도 구체적이다. 고정 사전분포로 프리픽스를 학습하는 Nested Dropout과 비교하면, PushT에서 Nested Dropout은 MixSIGReg 조합으로 96차원에서 92.50±0.76%가 최고였지만 ALeWM은 평균 53.49±10.51차원으로 96.00±0.00%를 냈다. OGBench-Cube에서는 Nested Dropout의 SIGReg 변형이 65.83~70.00%에 머문 반면 MixSIGReg 변형은 192차원에서 72.83±1.74%, ALeWM은 16차원만으로 79.00±0.76%를 기록했다. 고정 프리픽스 크기를 전부 스윕해도 모든 과제에서 균일하게 최적인 값은 없었고, ALeWM의 선택기는 OGBench-Cube에서 최적 고정값(79.00%, k=16)과 같거나 PushT(96.00% vs k=96의 95.83%), Reacher(85.83% vs k=128의 84.17%)에서 더 나으면서 평균 프리픽스는 더 짧았다. 정규화 항만 바꾼 비교에서는 MixSIGReg가 SIGReg 대비 PushT +2.50%포인트, OGBench-Cube +8.17%포인트를 얻었고 평균 선택 프리픽스는 각각 138.67→53.49, 128.00→16.00으로 줄었다.
실무 관점에서 이 논문은 잠재 공간 플래너를 만들 때 임베딩 폭을 줄이는 대신 정보 순서를 학습시키는 선택지를 제시한다. 다만 저자들이 밝힌 한계가 분명하다. 용량 네트워크를 월드모델과 공동 학습하기 때문에 초기화와 랜덤성에 따라 선택이 흔들리고, 시드별로 선택 용량이 달라질 수 있다(다만 각 데이터셋·백본에서 이웃한 지원 값에 머무는 경향). MixSIGReg 사전분포가 용량 순위에 대한 다항식과 그 차수 α에 의존한다는 점, straight-through Gumbel-Softmax의 그래디언트가 편향되어 있다는 점도 남은 문제로 꼽는다. 또한 장난감 실험 결과는 프리픽스 집중을 지지할 뿐 좌표 단위 분리나 실제 요인 개수 복원을 보장하지 않으며, 여러 데이터셋을 섞어 학습하면 데이터셋별로 다른 용량 분포를 학습하긴 하지만 성능은 떨어진다고 보고한다.