Skip to content

Latest commit

 

History

611 Commits

Folders and files

NameName
Last commit message
Last commit date
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 

Repository files navigation

TorchWM

PyPI version PyPI downloads License: MIT Documentation CI

Modular PyTorch library for world models — many algorithms, one consistent API.

TorchWM brings the major world-model families together under a single PyTorch API. Train Dreamer, PlaNet, JEPA, IRIS, DIAMOND, DiT, and Genie agents through create_config / create_model / make_env, or drop down to their encoders, decoders, and latent-dynamics backbones to compose your own architecture. Environment adapters (Gym/Gymnasium, DeepMind Control, MuJoCo, Brax, Atari, Unity ML-Agents) and ONNX / TorchScript / TensorRT export come built in.

Quick Start

# Install the core package from PyPI.
# This keeps environment integrations and experiment logging optional.
pip install torchwm

# With extras
pip install torchwm[gym]       # Gym/Gymnasium environments (runnable quick start)
pip install torchwm[dmc]       # DeepMind Control Suite (walker-walk, cheetah-run, ...)
                               # On CPython 3.13 also run: python -m torchwm.install_dmc
pip install torchwm[worldmodels] # Classic World Models (ConvVAE + CMA-ES controller)
pip install torchwm[ml-agents] # Unity ML-Agents
pip install torchwm[ml]        # TensorBoard, W&B logging
pip install torchwm[viz]       # Latent-space visualization (plotly, UMAP)
pip install torchwm[dev]       # Testing and linting

# Or add it to a uv-managed project.
uv add torchwm

TorchWM depends on PyTorch but does not force a single PyTorch wheel index. If you need a specific PyTorch build, install or add the PyTorch packages with the index recommended for your platform by the PyTorch installation selector:

# Example: CUDA 12.1 wheels. Choose a different index for CPU, ROCm, CUDA 11.x, CUDA 12.4+, or macOS.
uv add torch torchvision torchaudio --index https://download.pytorch.org/whl/cu121

Use the friendly top-level API for the common path. The example below runs on a base pip install torchwm[gym] — no simulator downloads required:

import torchwm

# Trains a Dreamer agent on a Gymnasium task. Bump `total_steps` for real runs.
# `seed_steps` of random play come first and count towards `total_steps`; the
# final checkpoint is written to `<logdir>/ckpts/` when training finishes.
agent = torchwm.create_model(
    "dreamer",
    env="Pendulum-v1",
    env_backend="gym",
    seed_steps=1_000,
    total_steps=10_000,
)
agent.train()

To train on DeepMind Control tasks such as walker-walk, install the DMC extra (pip install torchwm[dmc]) and use the default backend:

agent = torchwm.create_model("dreamer", env="walker-walk", total_steps=1_000_000)
agent.train()

Swap the algorithm, keep the code

Every algorithm in the table below is reachable through the same factory, so comparing them is a loop rather than a rewrite:

import torchwm

for algo in ["dreamer-v1", "dreamer-v2", "dreamer-v3"]:
    agent = torchwm.create_model(
        algo, env="Pendulum-v1", env_backend="gym", total_steps=20_000
    )
    agent.train()

examples/algorithm_comparison.py runs exactly this and writes a comparison plot. Construction is unified across all registered models. A shared step-budget train() currently covers the Dreamer family — other agents use their own trainers (torchwm train … / JEPAAgent.train() / DiamondAgent.train()). The example reports which is which rather than assuming.

Features

  • Unified interfaces across world-model algorithms
  • Modular encoders, decoders, dynamics models, and backbones
  • Training and inference utilities for model-based reinforcement learning
  • Environment integrations for Gym/Gymnasium, Unity ML-Agents, MuJoCo, Brax, and robotics extras
  • Optional logging, visualization, development, and documentation extras

Architecture

flowchart LR
    subgraph API["torchwm API"]
        CFG["create_config()"]
        MDL["create_model()"]
        ENV["make_env()"]
    end

    subgraph CONFIGS["Configs"]
        DC["DreamerConfig"]
        JC["JEPAConfig"]
        IC["IRISConfig"]
        GC["GenieConfig"]
        DIC["DiTConfig / DiamondConfig"]
    end

    subgraph AGENTS["Agents / Models"]
        DR["Dreamer / DreamerV1 / DreamerV2"]
        JP["JEPAAgent"]
        IR["IRISAgent"]
        GN["Genie"]
        DT["DiT / DIAMOND"]
    end

    subgraph BACKBONES["Backbones"]
        RSSM["RSSM / ModularRSSM"]
        VIT["VisionTransformer"]
        VQ["VQ-VAE / VideoTokenizer"]
        ST["STTransformer"]
        DIF["DDPM / DiT diffusion"]
    end

    subgraph ENVS["Environments"]
        GYM["Gym / Atari"]
        DMC["DeepMind Control"]
        MJ["MuJoCo"]
        BR["Brax"]
        UN["Unity ML-Agents"]
        ROB["Robotics"]
        more["..."]
    end

    subgraph EXPORT["Export"]
        ONNX["ONNX"]
        TS["TorchScript"]
        TRT["TensorRT"]
    end

    CFG --> CONFIGS
    MDL --> AGENTS
    ENV --> ENVS
    AGENTS --> BACKBONES
    AGENTS -.-> ENVS
    AGENTS --> EXPORT
Loading

Supported Algorithms

Every row is a registry entry — pass the name straight to torchwm.create_model(...) or torchwm.create_config(...). Run torchwm.list_models() for the live list.

Name Algorithm Description Key Features
dreamer Dreamer Model-based RL with latent dynamics (alias for dreamer-v1) Imagination, actor-critic
dreamer-v1 DreamerV1 Latent imagination with Gaussian heads Normal heads, standard KL
dreamer-v2 DreamerV2 Discrete latents for pixel control Symlog two-hot heads, balanced KL
dreamer-v3 DreamerV3 (name) Same DreamerAgent as dreamer Registry name for V3-style configs; not a separate paper-complete V3
planet PlaNet Latent planning from pixels, no explicit policy RSSM, CEM planner
modular-rssm ModularRSSM Composable recurrent state-space model Swappable priors/posteriors, custom heads
iris IRIS Sample-efficient RL with Transformers Discrete VAEs, world models
jepa JEPA Self-supervised visual representations Masked prediction, ViT
dit DiT Diffusion Transformer workflows Patch embeddings, diffusion backbones
diamond DIAMOND Diffusion world model for pixel-control RL EDM sampling, Atari imagination rollouts
genie Genie Generative interactive environments from video Latent actions, spatiotemporal transformer
genie-small Genie (small) Development- and test-sized Genie Same architecture, reduced width/depth
genie-large Genie (large) Scaled-up Genie variant Higher capacity dynamics + tokenizer

Documentation

Community

TorchWM follows semantic versioning as of 1.0.0. The public API — everything listed in the Public API reference and re-exported from the top-level torchwm namespace — will not break within the 1.x line; anything removed gets a deprecation warning for at least one minor release first. Submodule internals not listed there may still change.

About

A modular PyTorch library designed for learning, training, and deploying world models across various environments.

Topics

Resources

Code of conduct

Contributing

Security policy

Stars

29 stars

Watchers

1 watching

Forks

Releases

Packages

Used by

Contributors

Languages