LeWM vs PreJEPA vs GCRL: Which World Model Architecture Should You Choose?
Choose LeWM for deterministic model-predictive control with sampling-based planners, Pre-JEPA for multimodal JEPA-style representations with optional reconstruction, and GCRL when you need a stochastic goal-conditioned policy that outputs actions directly.
Selecting the right world model architecture in the stable-worldmodel repository depends on whether you need model-predictive control or a direct policy, multimodal inputs or pixels-only, and deterministic versus stochastic behavior. This guide compares LeWM, Pre-JEPA, and GCRL—three distinct implementations available in the GalilAI Group's open-source framework—to help you match the architecture to your robotics or control task.
LeWM: Deterministic Latent Dynamics for MPC
LeWM (LeWM class) implements a deterministic latent-state predictor designed for model-predictive control (MPC). The architecture encodes observations into embeddings and rolls out future states autoregressively using a learned dynamics model.
According to the source code in stable_worldmodel/wm/lewm/lewm.py, the core forward pass uses a rollout method that encodes initial observations and iteratively predicts next embeddings:
def rollout(self, info, action_sequence, history_size: int = 3):
# Encode initial observation → latent embedding
# Autoregressively predict next embedding with self.predict(...)
# Return info['predicted_emb'] (B × S × T × D)
Key characteristics of LeWM include:
- SIGReg regularization: The model applies a Gaussian regularizer to maintain a well-structured latent space during training.
- Criterion-based planning: The
criterionmethod returns a scalar cost for action candidates (MSE between predicted and goal embeddings), enabling compatibility with sampling-based solvers like CEM, iCEM, and MPPI. - Lightweight design: No decoder is required, reducing compute overhead compared to reconstruction-based architectures.
Training is handled by scripts/train/lewm.py, and the model is typically evaluated using planners that call model.get_cost(info, candidates) to evaluate action sequences.
Pre-JEPA: Multimodal Joint-Embedding Predictive Architecture
Pre-JEPA extends the JEPA (Joint-Embedding Predictive Architecture) paradigm with support for multimodal inputs and optional reconstruction. Implemented in stable_worldmodel/wm/prejepa/prejepa.py with helper modules in module.py, this architecture learns representations by predicting future patches from past patches plus auxiliary modalities.
The encode method demonstrates the multimodal capability:
def encode(self, info, pixels_key='pixels', ..., extra_encoders=None):
# Encode pixels → (B,T,P,D)
# For each extra key, run its encoder, tile across patches, concat
# Store result in info['emb'] (or custom target)
Pre-JEPA distinguishes itself through:
- Multiple extra encoders: You can inject proprioception, actions, or tactile data through dedicated
Embeddermodules that concatenate with visual patches. - Optional decoder: Unlike LeWM, Pre-JEPA supports reconstruction via a decoder (e.g.,
ViTDecoder), enabling visualization and downstream generative tasks. - JEPA training objective: Uses a "predict-the-future" contrastive-style loss without explicit reward modeling.
The API mirrors LeWM (encode, predict, rollout, criterion), making it compatible with the same MPC solvers while supporting richer sensory inputs. Training launches from scripts/train/prejepa.py.
GCRL: Goal-Conditioned Policy Learning
GCRL (GCRL class) takes a fundamentally different approach from the planning-based architectures above. Located in stable_worldmodel/wm/gcrl/gcrl.py, this architecture directly predicts a distribution over actions given observations and goals, functioning as a policy rather than a world model for planning.
The get_action method is the primary interface:
def get_action(self, info, sample=False, temperature=1.0):
# Encode observation & goal → latent embeddings
# Predict action means with self.predict_actions(...)
# Optional sampling using learned log_stds
Key implementation details include:
- Action distribution: A separate action predictor outputs means, while a learned log-std vector provides stochasticity for exploration.
- Direct policy execution: Unlike LeWM and Pre-JEPA, GCRL does not require external planners like CEM or MPPI; it outputs actions directly via
get_action. - RL training integration: Used by offline RL algorithms in
scripts/train/hilp.pyandscripts/train/gcivl.py, supporting value prediction and KL-regularization objectives.
GCRL is ideal when you need stochastic behavior for exploration or when you prefer imitation learning and offline RL over trajectory optimization.
Decision Framework: How to Select Your Architecture
When choosing between these three architectures in the stable-worldmodel framework, evaluate your requirements across five dimensions:
1. Control Paradigm: Planner vs Policy
- Planner (LeWM or Pre-JEPA): Choose when you need to search over candidate action sequences using
rolloutandcriterionmethods with solvers like CEM or MPPI. - Policy (GCRL): Choose when you want direct action inference via
get_actionwithout external trajectory optimization.
2. Input Modality
- Pixels only: Any architecture works.
- Multimodal (proprioception, actions, tactile): Pre-JEPA provides explicit support for extra encoders. GCRL also supports extra encoders but focuses on action prediction rather than latent state rollout.
3. Determinism vs Stochasticity
- Deterministic latent dynamics: LeWM provides pure deterministic rollouts ideal for consistent MPC.
- Stochastic policy: GCRL implements learned log-std parameters for sampling actions, necessary for exploration and uncertainty modeling.
4. Training Objective
- SIGReg + prediction loss: LeWM combines prediction loss with Gaussian regularization.
- JEPA + reconstruction: Pre-JEPA uses contrastive future prediction with optional decoder reconstruction.
- RL-style objective: GCRL supports value losses and KL-regularization as seen in GCIVL and GCIQL implementations.
5. Compute Budget
- Lightest: LeWM (no decoder, minimal overhead).
- Moderate: GCRL (adds log-std parameters and value networks).
- Heaviest: Pre-JEPA (includes decoder and multiple modality encoders).
Practical Implementation Examples
Loading LeWM for CEM Planning
import stable_worldmodel as swm
from stable_worldmodel.solver import CEMSolver
from stable_worldmodel.policy import WorldModelPolicy, PlanConfig
# Load checkpoint containing LeWM instance
model = swm.WM.load_checkpoint(
"path/to/lewm_checkpoint.pt",
model_cls=swm.wm.LeWM,
)
solver = CEMSolver(model=model, num_samples=300, horizon=10)
policy = WorldModelPolicy(solver=solver, config=PlanConfig(horizon=10))
world = swm.World("swm/PushT-v1", num_envs=8)
world.set_policy(policy)
results = world.evaluate(episodes=20)
Configuring Pre-JEPA with Proprioception
from stable_worldmodel.wm.prejepa import PreJEPA
from stable_worldmodel.wm.prejepa.module import Embedder, CausalPredictor
backbone = swm.models.ViT()
predictor = CausalPredictor(dim=512, num_layers=4)
action_enc = Embedder(in_dim=4, out_dim=64)
proprio_enc = Embedder(in_dim=10, out_dim=32)
decoder = swm.models.ViTDecoder()
model = PreJEPA(
encoder=backbone,
predictor=predictor,
extra_encoders=dict(action=action_enc, proprio=proprio_enc),
decoder=decoder,
history_size=3,
)
Direct Action Inference with GCRL
from stable_worldmodel.wm.gcrl import GCRL
encoder = swm.models.ViT()
action_predictor = swm.models.ActionPredictor(dim=512, out_dim=4)
value_predictor = swm.models.ValuePredictor(dim=512)
model = GCRL(
encoder=encoder,
action_predictor=action_predictor,
value_predictor=value_predictor,
history_size=3,
)
# Inference
obs = {"pixels": image_tensor, "goal": goal_image_tensor}
action = model.get_action(obs, sample=True, temperature=0.8)
Summary
- LeWM provides deterministic latent dynamics in
stable_worldmodel/wm/lewm/lewm.pyoptimized for sampling-based MPC with SIGReg regularization. - Pre-JEPA offers multimodal JEPA-style representation learning in
stable_worldmodel/wm/prejepa/prejepa.pywith optional reconstruction capabilities. - GCRL implements stochastic goal-conditioned policies in
stable_worldmodel/wm/gcrl/gcrl.pyfor direct action prediction without external planners. - Select based on your control paradigm (planning vs policy), modality requirements, and stochasticity needs.
Frequently Asked Questions
Can I use Pre-JEPA without the decoder to save compute?
Yes. According to the implementation in stable_worldmodel/wm/prejepa/prejepa.py, the decoder is optional. You can instantiate PreJEPA with decoder=None to use only the joint-embedding predictive architecture without reconstruction overhead, functioning similarly to LeWM but with multimodal encoder support.
Does GCRL support planning like LeWM and Pre-JEPA?
No. While GCRL includes encoders for observations and goals, it does not implement the rollout and criterion methods used by CEM and MPPI solvers. GCRL is designed for direct policy execution via get_action, making it unsuitable for trajectory optimization planners that require cost evaluation over action sequences.
Which architecture performs best for image-only observations?
All three architectures handle image-only inputs. LeWM is optimal for pure pixel-based MPC with minimal overhead. Pre-JEPA adds reconstruction capabilities if visualization is needed. GCRL is preferable when you want a direct policy rather than planning. For image-only tasks without multimodal sensors, LeWM offers the best performance-to-compute ratio.
How does the SIGReg regularizer in LeWM affect training?
The SIGReg regularizer in stable_worldmodel/wm/lewm/lewm.py constrains the latent space to follow a Gaussian distribution during training. This improves the stability of long-term rollouts in the rollout method by preventing latent drift, ensuring that autoregressive predictions remain within the training distribution when unrolled over long horizons for MPC.
Have a question about this repo?
These articles cover the highlights, but your codebase questions are specific. Give your agent direct access to the source. Share this with your agent to get started:
curl -s "https://instagit.com/install.md" Maintain an open-source project? Get it listed too →