Command Palette
Search for a command to run...
H-JEPA:視覚計画のための階層的世界モデルのエンドツーエンド学習
H-JEPA:視覚計画のための階層的世界モデルのエンドツーエンド学習
Wancong Zhang Basile Terver Michael Rabbat Yann LeCun Randall Balestriero
概要
潜在世界モデルを用いた長期計画には、複数の時間スケールと抽象度にわたる推論が必要である。既存のタスク非依存型JEPA世界モデルは、単一の潜在空間(多くの場合、単一の時間スケール)で予測と計画を行う。我々は、各階層が自身の学習済み潜在空間においてより遠い未来を予測する、行動条件付きJEPAの階層を訓練するためのエンドツーエンドの手法であるH-JEPAを紹介する。計画はトップダウンで進行する。最上位層が目標への進捗を最適化し、各層の予測がその下のプランナーにとってのサブゴールとなる。データ内の要因が分離した時間スケールで発展する場合、上位層は速く変化する予測不能な詳細を捨て、遅く変化するタスク関連の状態を保持する。4つのシミュレーション環境(ナビゲーションと操作)において、階層的計画はフラットなJEPAよりも優れている。Visual AntMazeでは、3層の階層により、プランナーの計算量を抑えつつ成功率が18%から73%に向上した。アブレーション研究により、これらの利点は時間的分解と上位層の目標表現の両方に起因することが示された。逆力学の教師信号を用いることで、この手法はDROIDからの多様な実ロボットビデオにも拡張され、階層により低いプランナー計算量でオフライン計画の忠実度が向上する。
One-sentence Summary
H-JEPA, developed by researchers at NYU, Advanced Machine Intelligence, INRIA Paris, and Brown University, introduces an end-to-end recipe for training a hierarchy of action-conditioned JEPA world models where each level predicts farther ahead in its own latent space, enabling top-down planning with subgoals; on Visual AntMaze, a three-level hierarchy raises success from 18% to 73% with less planner compute, and with inverse-dynamics supervision the approach extends to real-robot DROID videos.
Key Contributions
- Introduces H-JEPA, an end-to-end training method for a hierarchy of action-conditioned JEPA world models, where each level predicts in its own latent space at progressively longer temporal strides and higher abstraction, without reward or reconstruction objectives.
- Demonstrates that hierarchical planning outperforms a flat JEPA baseline across four simulated navigation and manipulation environments; on Visual AntMaze, a three-level hierarchy raises success from 18% to 73% using less planner compute.
- Identifies two complementary mechanisms behind these gains: temporal decomposition into shorter subgoals and higher-level goal representations that score progress in more abstract spaces; adding an inverse-dynamics term extends the method to diverse real-robot videos from DROID, improving offline planning fidelity at lower planner compute.
Introduction
World models let agents learn environment dynamics from experience so they can predict and plan, and recent Joint-Embedding Predictive Architectures (JEPAs) do this by forecasting future latent states instead of pixels or rewards. Yet these models typically operate in a single latent space at a single timescale, which creates two problems: long-horizon prediction requires many fine-grained rollout steps, compounding errors and enlarging the action search space, and a shared latent space must simultaneously capture low-level dynamics and support abstract goal matching, often failing at the latter when goals are conceptual (like reaching a location) rather than state-specific. Prior hierarchical approaches either reconstruct pixels without planning or rely on reward-driven, task-specific policies; the closest task-agnostic model, HWM, plans across multiple horizons but still uses one latent space, so it cannot match objectives at higher levels of abstraction.
The authors introduce H-JEPA, a hierarchical JEPA architecture where each level predicts future states in its own latent space, using progressively coarser temporal strides and more abstract representations, with no reward or reconstruction objectives. Their key contributions are an end-to-end training method for hierarchical JEPA models, a demonstration that higher levels discard unpredictable fast details while retaining slower predictable state, and evidence that hierarchical planning outperforms single-level planning at lower compute by enabling goal scoring in more abstract spaces and decomposing long-horizon tasks into easier subproblems. They also show that adding an inverse-dynamics term extends the method to real-robot manipulation data from DROID, where scene, lighting, and objects change across episodes, improving offline planning fidelity despite such variability.
Dataset
The authors evaluate their hierarchical model, H-JEPA, across several environments that span navigation and manipulation tasks. The dataset composition and usage are as follows:
-
Environments and Sources
- Navigation: FourRoomDistractors, which pairs a controlled agent with a distractor that moves smoothly but teleports randomly, and Visual AntMaze, where a quadruped navigates a maze.
- Manipulation: Push-T (pushing a T-shaped block) and OGBench Cube (picking up and placing a cube).
- Real-world extension: DROID, a dataset of real-robot manipulation trajectories, used for additional evaluation (§4.3).
-
Training Setup and Data Processing
- The model is trained end-to-end with up to four hierarchical levels.
- The level-1 JEPA uses a ViT-Tiny encoder with a [CLS] token to summarize each image globally, plus a causal transformer predictor. This single-level model serves as the flat baseline.
- Higher-level JEPAs use two-layer MLP encoders and causal transformer predictors.
- At each higher level, a stride of sℓ=2 and window wℓ=1 mean that one prediction step at that level spans twice as many environment timesteps as the level below.
- All levels in a given environment share the same SIGReg coefficient.
- No explicit data augmentation or filtering is mentioned; the datasets are used as provided.
-
Evaluation Protocol (Data Usage in Probing)
- The authors train two-layer MLP probes on frozen representations from each level.
- They report variance-normalized MSE (NMSE), averaged equally over an entity's dimensions. An NMSE near 1 means predicting the marginal mean, while near 0 indicates accurate recovery.
- Results are mean ± SE over three training seeds.
- Pixel decodings are qualitative examples from post-hoc decoders.
-
Dataset Statistics and Observed Patterns
- Higher levels discard fast-varying features while retaining slow ones. For example, in AntMaze, the body state (joints) becomes less recoverable with depth, while the global position remains accurate. In FourRoom, the distractor position becomes less recoverable because longer prediction horizons often cross random teleports.
- The paper relates per-entity probing error at level 2 to the "entity frequency gap" in the training data: the ratio of spectral centroids of the fastest- and slowest-varying state components.
- Larger frequency gaps are associated with greater loss of fast-entity information at level 2, while slow-entity recoverability stays near its level-1 value.
- This selective abstraction appears in AntMaze, Humanoid, and FourRoom, but not in the manipulation datasets (Push-T, OGBench Cube), which have smaller frequency gaps.
- Gap measurements are detailed in §F.2.
-
Additional Processing Notes
- The paper does not describe any cropping strategy or manual metadata construction for these datasets; the inputs are raw observations (images for navigation/manipulation, plus trajectory data for DROID).
- The model's hierarchical abstractions emerge from joint optimization of encoder and predictor, where prediction-error gradients encourage each level to retain only features predictable at its own timescale.
Method
H-JEPA learns a hierarchy of latent predictive models from trajectories of observations and actions. The hierarchy is temporal: level 1 operates on the finest observation stream, while each higher level consumes the latent states produced by the level below over a coarser stride. This gives a stack of JEPA world models in which each level has its own observation encoder, action encoder, and latent predictor. All levels predict future latent states from past latent states and actions.
Let ot denote an observation and at the action block for the transition from ot to ot+1. The first level encodes the observation, optionally with the proprioceptive state, into a latent state zt(1)=E(1)(ot), and the action block into at(1)=A(1)(at). Higher levels are built compositionally: both state and action encoders pool temporal windows of latents from the level below.
Each level ℓ>1 has two temporal hyperparameters. The stride sℓ is the subsampling factor: upper-level time t maps to lower-level time t⋅sℓ, so consecutive upper-level states are sℓ lower-level steps apart. The window size wℓ is the number of lower-level steps each upper-level state summarizes. The two are independent: sℓ sets how far each window advances and wℓ how far it spans.
With a:b denoting the b−a indices a,…,b−1, the level-ℓ state and action embeddings are:
zt(ℓ)=E(ℓ)(zt⋅sℓ:t⋅sℓ+wℓ(ℓ−1)),at(ℓ)=A(ℓ)(at⋅sℓ+wℓ−1:(t+1)⋅sℓ+wℓ−1(ℓ−1)).The state encoder E(ℓ) pools the wℓ-step lower-level window into one abstract state. The action encoder A(ℓ) instead aggregates sℓ lower-level action embeddings, independent of wℓ, so at(ℓ) covers the full transition from the state anchoring zt(ℓ) toward the next upper-level state zt+1(ℓ). All experiments use wℓ=1, so upper-level state encoders are pointwise.
Each level is trained as a JEPA in its own latent space. Given a context of cℓ latent states (zt−cℓ+1(ℓ),…,zt(ℓ)) and the associated action embeddings, a predictor F(ℓ) predicts future latent states. One-step training uses the teacher-forced latent prediction loss:
Lpred(ℓ)=cℓ1τ=1∑cℓF(ℓ)(z<τ(ℓ),a<τ(ℓ))−zτ+1(ℓ)22.The per-level loss also includes SIGReg, a sketched normality regularizer that prevents collapse by encouraging an isotropic Gaussian embedding distribution. Let Z(ℓ)∈Rn×D collect n level-ℓ embeddings of width D, flattened over batch and time. The overall level-ℓ objective is:
L(ℓ)=Lpred(ℓ)+λℓSIGReg(Z(ℓ)).H-JEPA plans top-down through the learned hierarchy. The current and goal observations are encoded at every level ℓ, yielding initial latent states {z0(ℓ)}ℓ=1L and goal latent states {g(ℓ)}ℓ=1L, where L is the number of levels. The top level proposes a coarse plan to the goal. Each lower level plans toward subgoals supplied by the predicted trajectory of the level above, refining the plan into finer-scale transitions until level 1 produces primitive actions. The top-level planner optimizes macro-actions to minimize the distance between its final predicted state and the goal:
a0:HL−1(L),∗=arga0:HL−1(L)minz^HL(L)−g(L)22.At each level ℓ, the predicted states (z^1(ℓ),…,z^Hℓ(ℓ))=Roll(ℓ)(z0(ℓ),a0:Hℓ−1(ℓ)) are obtained by autoregressively applying predictor F(ℓ) from z0(ℓ) under the candidate action sequence, with horizon Hℓ measured in level-ℓ steps. A superscript * marks the optimized actions and their predicted states.
For each lower level ℓ<L, the optimized rollout from the level above, (z^1(ℓ+1),∗,…,z^Hℓ+1(ℓ+1),∗), supplies subgoals. The lower level optimizes its actions to match these subgoals with its own predicted states, encoded into the upper-level latent space by E(ℓ+1).
the paper first choose the number of subgoals to match, 1≤Kℓ+1≤Hℓ+1, then set the lower-level horizon to cover the corresponding encoder windows. With upper-level stride sℓ+1 and window size wℓ+1, this requires Hℓ=Kℓ+1sℓ+1+wℓ+1−1. The encoded predictions are:
z~i(ℓ+1)=E(ℓ+1)(z^isℓ+1:isℓ+1+wℓ+1(ℓ)),i=1,…,Kℓ+1.Each z~i(ℓ+1) is aligned with its upper-level subgoal in time and latent space. The lower level solves:
a0:Hℓ−1(ℓ),∗=arga0:Hℓ−1(ℓ)min[z~Kℓ+1(ℓ+1)−z^Kℓ+1(ℓ+1),∗22+βi=1∑Kℓ+1−1z~i(ℓ+1)−z^i(ℓ+1),∗22],where β≥0 weights each intermediate subgoal relative to the terminal target. Closed-loop planning uses Kℓ+1=1, matching only the first predicted subgoal; open-loop planning uses Kℓ+1=Hℓ+1, with β=0 matching only the terminal prediction and β>0 also matching intermediate subgoals.
the paper optimize action sequences at every level by gradient descent. Level 1 returns primitive actions for execution in the environment. In closed-loop control, the paper execute an action prefix, encode the new observation, and repeat hierarchical planning until the evaluation budget of environment steps is exhausted.
Experiment
Experiments evaluated hierarchical JEPA (H-JEPA) models across navigation and manipulation tasks, including real-robot DROID data. Learning hierarchical abstractions showed that higher levels discard fast-varying features while retaining slow ones, but this selective abstraction only emerges when entity frequency gaps are large, not in manipulation datasets with smaller gaps. In planning, H-JEPA improved success-compute tradeoffs over flat baselines and a single-latent-space hierarchical model, with gains attributed to both more informative abstract cost geometries and temporal decomposition into easier subgoals. On diverse scenes in DROID, adding an inverse-dynamics loss was necessary to avoid slow-feature collapse, after which H-JEPA outperformed flat and hierarchical baselines at lower planning budgets.
On AntMaze, replacing the native level-1 planning cost with costs computed in higher-level latent spaces of the same H-JEPA model consistently improves flat-planning success across model depths, with the largest gains seen when native-space planning is weakest. Hierarchical planning that also decomposes search temporally achieves substantially higher success than any cost-only variant, indicating that both cost geometry and temporal abstraction contribute to planning performance. For every multilevel H-JEPA tested, at least one upper-level cost improves success over the native level-1 cost. The best projected level-1 planner exceeds the separate LeWM baseline in mean success at every tested depth. Gains from changing only the cost are largest where native level-1 planning performs poorly. Hierarchical planning, which changes both cost and temporal decomposition, clearly outperforms cost-only variants.
On AntMaze, evaluating different planning configurations for H-JEPA models showed that using higher-level latent costs instead of the native level-1 cost consistently boosts flat-planning success, with the biggest improvements where native planning struggles. Hierarchical planning that additionally decomposes search over time outperforms all cost-only variants, confirming that both cost geometry and temporal abstraction matter. Across every multilevel H-JEPA, at least one upper-level cost improves success over the baseline, and the best projected level-1 planner surpasses the separate LeWM baseline at all tested depths. Overall, cost modification alone yields the largest gains when native planning is weakest, while full hierarchical planning delivers the strongest results.