Command Palette
Search for a command to run...
H-JEPA: END-TO-END-LERNEN HIERARCHISCHER WELTMODELLE FÜR VISUELLE PLANUNG
H-JEPA: END-TO-END-LERNEN HIERARCHISCHER WELTMODELLE FÜR VISUELLE PLANUNG
Wancong Zhang Basile Terver Michael Rabbat Yann LeCun Randall Balestriero
Zusammenfassung
Langfristige Planung mit latenten Weltmodellen erfordert Schlussfolgerungen über mehrere Zeitskalen und Abstraktionsebenen hinweg. Bestehende aufgabenagnostische JEPA-Weltmodelle sagen vorher und planen in einem einzigen latenten Raum, oft auf einer einzigen Zeitskala. Wir stellen H-JEPA vor, ein End-to-End-Verfahren zum Training einer Hierarchie aktionskonditionierter JEPAs, bei dem jede Ebene weiter in die Zukunft in ihrem eigenen gelernten latenten Raum vorhersagt. Die Planung erfolgt von oben nach unten: Die oberste Ebene optimiert den Fortschritt in Richtung des Ziels, und die Vorhersagen jeder Ebene werden zu Teilzielen für den darunterliegenden Planer. Wenn sich Faktoren in den Daten auf getrennten Zeitskalen entwickeln, verwerfen höhere Ebenen schnelle, unvorhersehbare Details und behalten langsamere aufgabenrelevante Zustände bei. In vier simulierten Navigationsund Manipulationsumgebungen verbessert die hierarchische Planung einen flachen JEPA; im visuellen AntMaze erhöht eine dreistufige Hierarchie den Erfolg von 18 % auf 73 % bei geringerem Planungsaufwand. Ablationsstudien führen diese Gewinne sowohl auf die zeitliche Zerlegung als auch auf die Zielrepräsentationen der höheren Ebenen zurück. Mit inverser Dynamiküberwachung erweitert sich der Ansatz auf vielfältige reale Roboter-Videos aus DROID, wo die Hierarchie die Offline-Planungstreue bei geringerem Planungsaufwand verbessert.
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.