SIGN IN SIGN UP

feat: add LeWM, an action-conditioned latent world model trained with SIGReg (#2032)

* Add LeWM latent world model

Predictor, action encoder, loss and a PyTorch example, scoped to what LeWM
alone needs. Later world models add arguments that default to this behavior,
so nothing here changes meaning when they land.

Co-Authored-By: Claude Opus 5 <noreply@anthropic.com>

* Reach the fused attention kernel through getattr

Type checking runs against the oldest supported torch, where
scaled_dot_product_attention does not exist yet.

Co-Authored-By: Claude Opus 5 <noreply@anthropic.com>

* refactor: require fused attention in LatentDynamicsPredictor

Co-Authored-By: Claude Opus 4.8 <noreply@anthropic.com>

* docs: document keyword-only forward and ONNX export in LeWMLoss

Co-Authored-By: Claude Opus 4.8 <noreply@anthropic.com>

* feat: add LeWMProjectionHead to the head api

Co-Authored-By: Claude Opus 4.8 <noreply@anthropic.com>

* refactor: parameterize LeWM example and use LeWMProjectionHead

Co-Authored-By: Claude Opus 4.8 <noreply@anthropic.com>

* refactor: remove unnecessary forward comment in LeWMLoss

Co-Authored-By: Claude Opus 4.8 <noreply@anthropic.com>

* refactor: rely on LeWM example module defaults instead of constants

Co-Authored-By: Claude Opus 4.8 <noreply@anthropic.com>

* feat: make LeWM predictor conditioning optional and expose PredictorBlock

Add conditional and causal flags to LatentDynamicsPredictor for actionless and bidirectional predictors, and export the AdaLN block as PredictorBlock. Skip predictor tests when torch lacks scaled_dot_product_attention (fixes minimal-deps CI).

Co-Authored-By: Claude Opus 4.8 <noreply@anthropic.com>

* feat: add batch_norm opt-out to LeWMProjectionHead

Mirror the batch_norm flag of the other projection heads so LeWM can drop the BatchNorm that diverges between train and eval on the rollout path.

Co-Authored-By: Claude Opus 4.8 <noreply@anthropic.com>

* docs: scope LeWM loss docstrings to continuous-latent methods

Rescope the 'every latent world model' claims in latent_distance and LeWMLoss to the continuous-latent family and mark them experimental; note the same on the LeWM example page.

Co-Authored-By: Claude Opus 4.8 <noreply@anthropic.com>

* fix: validate LeWM loss embeddings and guard rollout output_dim

Reject embeddings whose batch or width differ from predicted in LeWMLoss, and raise in rollout when output_dim != input_dim with steps > 1. Train the predictor in the causality test so it exercises attention.

Co-Authored-By: Claude Opus 4.8 <noreply@anthropic.com>

* docs: fix LeWM example shape contract, timm scope and run command

Clarify that the predictor returns only predicted embeddings, scope the timm requirement to the example, and correct the run path to examples/pytorch/lewm.py.

Co-Authored-By: Claude Opus 4.8 <noreply@anthropic.com>

* refactor: extract AdaLNZero conditioning module in LeWM predictor

Co-Authored-By: Claude Opus 4.8 <noreply@anthropic.com>

---------

Co-authored-by: Claude Opus 5 <noreply@anthropic.com>
F
fruet committed
c2db8f561ec653ab1f426ced9b44d3b10578a206
Parent: 3684a18
Committed by GitHub <noreply@github.com> on 9/28/2026, 6:22:57 PM