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