# LeWM vs PreJEPA vs GCRL: Which World Model Architecture Should You Choose?

> Compare LeWM, PreJEPA, and GCRL world models. Select LeWM for MPC, PreJEPA for multimodal JEPA, and GCRL for stochastic goal-conditioned policies. Find the best fit for your project.

- Repository: [GalilAI-group/stable-worldmodel](https://github.com/galilai-group/stable-worldmodel)
- Tags: comparison
- Published: 2026-05-30

---

**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`](https://github.com/galilai-group/stable-worldmodel/blob/main/stable_worldmodel/wm/lewm/lewm.py), the core forward pass uses a `rollout` method that encodes initial observations and iteratively predicts next embeddings:

```python
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 `criterion` method 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`](https://github.com/galilai-group/stable-worldmodel/blob/main/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`](https://github.com/galilai-group/stable-worldmodel/blob/main/stable_worldmodel/wm/prejepa/prejepa.py) with helper modules in [`module.py`](https://github.com/galilai-group/stable-worldmodel/blob/main/module.py), this architecture learns representations by predicting future patches from past patches plus auxiliary modalities.

The `encode` method demonstrates the multimodal capability:

```python
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 `Embedder` modules 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`](https://github.com/galilai-group/stable-worldmodel/blob/main/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`](https://github.com/galilai-group/stable-worldmodel/blob/main/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:

```python
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.py`](https://github.com/galilai-group/stable-worldmodel/blob/main/scripts/train/hilp.py) and [`scripts/train/gcivl.py`](https://github.com/galilai-group/stable-worldmodel/blob/main/scripts/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 `rollout` and `criterion` methods with solvers like CEM or MPPI.
- **Policy (GCRL)**: Choose when you want direct action inference via `get_action` without 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

```python
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

```python
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

```python
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.py`](https://github.com/galilai-group/stable-worldmodel/blob/main/stable_worldmodel/wm/lewm/lewm.py) optimized for sampling-based MPC with SIGReg regularization.
- **Pre-JEPA** offers multimodal JEPA-style representation learning in [`stable_worldmodel/wm/prejepa/prejepa.py`](https://github.com/galilai-group/stable-worldmodel/blob/main/stable_worldmodel/wm/prejepa/prejepa.py) with optional reconstruction capabilities.
- **GCRL** implements stochastic goal-conditioned policies in [`stable_worldmodel/wm/gcrl/gcrl.py`](https://github.com/galilai-group/stable-worldmodel/blob/main/stable_worldmodel/wm/gcrl/gcrl.py) for 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`](https://github.com/galilai-group/stable-worldmodel/blob/main/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`](https://github.com/galilai-group/stable-worldmodel/blob/main/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.