Multi-Modal

[2026-1] 정인아 - Learning Invariant Visual Representations for Planning with Joint-Embedding Predictive World Models

kino(키노) 2026. 8. 23. 17:07

논문 제목 : Learning Invariant Visual Representations for Planning with Joint-Embedding Predictive World Models (Bis-JEPA)

논문 링크 : https://arxiv.org/pdf/2602.18639

 

Introduction

  • generative world model에서 JEPA 계열이 주목받고 있다.(2026)
    • 이는 raw observation을 reconstrction하지 않고 latent representation을 prediction하는 방식이다.
  • Bis-JEPA는 base로 DINO-WM을 둔다.
    • DINO-WM는 JEPA 아이디어를 가지고, DINOv2 feature를 사용해 latent dynamics를 학습한 연구다.
    • reward 없이도 새로운 task에서 zero-shot planning이 가능함을 보였다.
  • 그러나, JEPA가 latent representation이 실제 task dynamics보다 slow feature에 더 많이 의존한다는 문제가 드러났다.
    • 여기서 slow feature는 planning에 필요 없는 visual feature를 뜻한다. (ex. background color, lighting..)
    • 그 이유는, 이 slow feature들이 시간에 따라 잘 안 변하기 때문에 JEPA의 predictive objective를 쉽게 만족할 수 있기 때문이다.
    • 이는, 단순히 좋은 visual encoder를 쓰는 것으로 해결되지는 않는데, 대부분 foundation encoder가 control-irrelevant 정보까지 함께 인코딩하기 때문이다.
  • 즉, Bis-JEPA는 JEPA의 latent predictive 장점은 유지하면서 slow feature에 robust한 world model을 만드는 것을 목표로 한다.
  • 그럼 어떻게 slow feature에 robust하게 만들 수 있을까?
    • 이들은 On-Policy Bisimulation Metric 개념을 사용한다. bisimulation이란 두 state가 같은 transition dynamics를 가지면 같은 state로 취급하자는 것이다.
    • 즉, control에 영향을 주지 않는 정보는 latent에서 제거하고, transition behavior가 비슷한 state는 가까이 두는 것을 제안한다.

 

Preliminaries

(1) JEPA 기본 구조

  • observation o_t를 latent z_t로 바꾼다.
    • $z_t=f_\theta(o_t)$
  • 그리고 z_t와 action a_t로 다음 latent \hat z_{t+1}를 예측한다.
    • $\hat z_{t+1}=\tilde T_\phi(z_t,a_t)$
  • 정답은 실제 다음 observation을 같은 encoder에 넣은 z_{t+1}를 사용한다.
    • $z_{t+1}=f_\theta(o_{t+1})$

(prediction loss)

$$\mathcal L_{\mathrm{pred}} =\mathbb E\left[\ell\left(\tilde T_\phi(f_\theta(o_t),a_t),f_\theta(o_{t+1})\right)\right].$$

 

(2) Bisimulation

  • 이때 Bis-JEPA는 free reward가 목표기 때문에 reward term을 버린다.
  • 따라서, 비슷한 transition이면 같은 state로 보는 것을 목표로 한다.

$$d_\pi(z,z')\approx\gamma W_1\left(P_\pi(\cdot|z),P_\pi(\cdot|z')\right)$$

  • Bis-JEPA는 visual encoder의 $z_t$를 그대로 쓰지 않고 $w_t=h_\eta(z_t)$ 라는 새로운 latent를 쓴다.
    • $o_t → z_t → w_t → \hat w_{t+1}$. $w_t$는 control-relevant latent representation이 된다.
    • dynamics도 w-space에서 학습한다.
    • $\hat w_{t+1} = T_\phi(w_t,a_t).$

 

(3) Loss function

  • Dynamics loss (prediction loss)
    • 다음 latent를 잘 예측하도록 만드는 loss

  • Bisimulation loss
    • 두 state w, w’의 latent distance가 미래 dynamics distance와 비슷하도록 만든다.
    • $w_t\approx w_t’$

 

Method

  • 구체적으로 이들은 pretrained visual encoder 뒤에 bisimulation encoder를 붙였다.
  • visual encoder로는 DINOv2, SimDINOv2, iBOT를 쓴다.
    • $f_\theta:\mathcal O\rightarrow\mathbb R^{N_p\times d_z}$ 형태의 patch feature를 만든다. 만든 이 feature $z_t$는 control-irrelevant 정보(slow feature)까지 포함한다.
  • bisimlation encoder는 $h_\eta : \mathbb R^{N_p\times d_z}\rightarrow\mathbb R^{N_p\times d_w}$ 를 사용해서, visual encoder의 각 patch의 dimensions을 줄인다.
  • 이 과정에서 representation collapse 문제가 생기는데, 모든 state $w$=0으로 보낼 경우, pairwise distance가 모두 0으로 되어 bisimulation loss를 쉽게 줄일 수 있다는 점이다.
    • 따라서 저자들은 VICReg를 사용해서 collapse하지 않도록 dimension에 variance를 유지시킨다.
      • 어떻게? Var(w)가 너무 작아지면 penalty를 주는 식이다.
  • 이때 표준 VICReg도 문제가 발생하는데, visual encoder의 feature에서 큰 variance를 가진 부분이 task-relevant dynamics 보다 slow feature와 더 강하게 relevant되어 있었기 때문이다.
    • 따라서, 표준 VICReg는 기본적으로 모든 representation 방향에 동일한 variance를 유지하기 때문에, slow feature들도 보존하게 되어 이들을 제거하기 어려웠음.
    • 저자들은 PCA-based VCReg loss를 제안한다.

  • $w_t$를 얻은 transformer가 dynamics를 학습한다.
    • $(w_{t-H:t-1},a_{t-H:t-1}) → \hat w_t.$
    • 현재 observation을 encode하고 후보 action sequence를 넣어 rollout한다.
      • $w_0 = h(f(o_0))$
      • $w_0 → \hat w_1 → … → \hat w_T$
    • goal observation w_g = h(f(o_g)) 와 가까워지는 action sequence를 찾는다.
      • $c = \|\hat w_T-w_g\|_2^2.$
      • predicted 최종 latent가 goal latent와 얼마나 가까운가를 cost로 사용하기 때문에 reward 필요없이 학습시킬 수 있다.

 

Experiments

  • setting
    • MuJoCo의 PointMaze를 사용
    • 공이 x,y 평면에서 goal까지 이동하는 navigation task
    • 모든 모델은 2,000개의 random trajectory로 학습, 50개의 random initial-goal pair 평가
    • 평가 metric은 success rate
  • baseline
    • DINO-WM
    • DINO-WM + domain randomization
    • DINO-Bisim (Ours)
  • env condition 
    • NC (변화 없음)
    • SC (약한 background change)
    • C (color/gradient change)
    • LC (큰 color change)
    • LCG (큰 color + gradient change)
    • D (moving distractor)

  • results

Success rate of DINO-Bisim under different scenarios compared to DINO-WM with and without DR.
Success rate with different pretrained visual encoder back- bones under the different scenarios.

 

Conclusion

  • Bis-JEPA는 latent representation에서 control-irrelevant한 정보는 제거하고, relevant한 state equivalence를 강제하는 bisimulation encoder를 제안했다.
  • 이는 reward prediction 없이도 invariant representation 학습이 가능하다는 점을 보여준다.