Skip to content

About

Diffusion-augmented MDPs in PyTorch: DA:PPO, DA:REPPO, DA:WPO and reproducible multimodal robotics demos

Topics

Resources

Stars

0 stars

Watchers

0 watching

Forks

Repository files navigation

DA-MDP PyTorch

Diffusion policies for maximum-entropy reinforcement learning: DA:PPO, DA:REPPO, and DA:WPO, with a shared PyTorch training loop. Paper · Multimodal demos and configs · Checkpoints

For talks and posters: repository QR code (SVG) or PNG.

Install the environment

The tested setup is Linux, Python 3.12, and PyTorch 2.9.1. The small built-in environments run on CPU. The robotics demos require an NVIDIA GPU and a driver compatible with the locked CUDA 12.8 PyTorch wheels.

  1. Install uv:

    curl -LsSf https://astral.sh/uv/install.sh | sh

    Open a new terminal if uv is not yet on your PATH.

  2. Clone the repository and create the locked environment:

    git clone https://github.com/Atarilab/DA_MDP_pytorch.git
    cd DA_MDP_pytorch
    uv sync --locked

    uv installs Python 3.12 if needed and creates .venv. Use uv run for commands below; activating the environment manually is unnecessary.

  3. Run a small CPU training check:

    uv run --locked python examples/train.py \
      --env turning_double_well --algo da_mdp_ppo \
      --device cpu --num-envs 8 --max-iters 1 --log-dir none

    This checks sampling and an optimizer update; one iteration is not a converged policy.

Robotics dependencies

For StackCube and PushT, install the ManiSkill extra:

uv sync --locked --extra maniskill3
uv run --locked --extra maniskill3 python -c \
  "import torch, mani_skill; print('CUDA available:', torch.cuda.is_available())"

The demo commands require CUDA available: True. ManiSkill/SAPIEN uses Vulkan for rendering; use a working NVIDIA driver with Vulkan support, including on headless machines. The first GPU simulation can download SAPIEN's PhysX GPU library. See ManiSkill installation for driver and renderer setup.

For MuJoCo Playground, use:

uv sync --locked --extra mujoco-playground
uv run --locked --extra mujoco-playground python examples/train.py \
  --config config/checks/mujoco_playground_pendulum_swingup_da_mdp_ppo.yaml \
  --device cuda:0 --max-iters 1

The learner is PyTorch; this optional simulator uses JAX/MJX/Warp. For both simulator backends, pass both extras to uv sync and uv run.

For IsaacLab, first install its supported environment using the IsaacLab instructions, then run python -m pip install -e . with that environment's Python. Run python examples/train.py --config config/checks/isaaclab_ant_da_mdp_ppo.yaml --device cuda:0 --max-iters 1 inside that environment. Its dependency versions are managed by IsaacLab rather than this repository's uv.lock.

Multimodal behavior

Both examples use one DA:REPPO policy and the same initial environment state. Changing the diffusion-prior sample produces different successful strategies.

StackCube: two stacking orders PushT: two rotation directions
DA:REPPO stacks blue on red and red on blue DA:REPPO rotates the T clockwise and counter-clockwise
Checkpoint-matched config Checkpoint-matched config

These are selected successful rollouts illustrating distinct modes. They are not aggregate success-rate measurements. Both GIFs were replayed from the released checkpoints using this public implementation.

Replay the demos for checkpoint downloads, exact commands, policy seeds, sampler settings, and GIF generation. Checkpoints are release assets; no training is needed to replay them.

Training configurations

Environment DA:PPO DA:REPPO DA:WPO
MuJoCo Playground Reference Reference Reference
G1 locomotion K=8 K=8 K=8
StackCube Paper Paper Paper
PushT Experimental Paper Experimental

G1 also includes K=2 and K=4 configurations. StackCube includes a REPPO action-chunking configuration. Paper presets preserve the experiment settings and use complete training budgets, including continuation stages. Current code includes subsequent correctness fixes; see configuration provenance and evaluation.

The main continuous-control curves in the paper use JAX. The generic MuJoCo presets here are PyTorch references, not exact reproductions of those curves. PushT PPO/WPO are runnable comparison presets without established learning results. Installation checks do not establish policy performance.

Train StackCube DA:REPPO with local TensorBoard logging:

uv run --locked --extra maniskill3 python examples/train.py \
  --config config/maniskill/stackcube/da_reppo.yaml \
  --device cuda:0 --seed 1 --output-dir outputs/stackcube

For a short installation check, append --num-envs 4 --num-steps-per-env 2 --num-mini-batches 1 --max-iters 1. These overrides change the training batch; remove them for the documented experiment. Full manipulation presets use 1,024 environments and substantial GPU memory.

For MuJoCo, select the corresponding config and task:

uv run --locked --extra mujoco-playground python examples/train.py \
  --config config/mujoco/da_reppo.yaml --task-id CheetahRun \
  --device cuda:0 --output-dir outputs/cheetah

Evaluate a manipulation checkpoint on a fixed initial state:

uv run --locked --extra maniskill3 python examples/evaluate_maniskill3_trudi.py \
  --config config/maniskill/stackcube/da_reppo.yaml \
  --checkpoint outputs/stackcube/model_final.pt \
  --task-id StackCubeSymmetricTest-v1 --sampling ode \
  --num-rollouts 128 --output outputs/stackcube/fixed_state.json

Replace the checkpoint path with the actual saved filename. This fixed-state check is distinct from the paper's averages over many IID initial states. See the evaluation protocols.

The released GIF checkpoints use their own demo configurations, including the older StackCube observation layout. Keep those configurations for replay; use the table above for new training.

Algorithms and development

examples/train.py accepts --algo da_mdp_ppo, da_mdp_reppo, or da_mdp_wpo for the built-in CPU checks. With --config, the YAML selects the algorithm. Conventional PPO/REPPO and action-space maximum-entropy WPO baselines are available under config/baselines/.

See the algorithm documentation for objectives, score guidance, diffusion samplers, and action chunking.

uv run --locked pytest -q
uv run --locked ruff check --select E4,E7,E9,F rsl_rl examples tools tests

Release validation records the checks and their limits.

Paper and attribution

Please cite the DA-MDP paper:

@article{sanokowski2025diffusionaugmented,
  title={Diffusion-Augmented Markov Decision Processes for Maximum Entropy Reinforcement Learning},
  author={Sanokowski, Sebastian and Patil, Kaustubh},
  journal={arXiv preprint arXiv:2512.02019},
  year={2025}
}

The shared training core derives from RSL-RL. Its BSD-3-Clause license, source copyright notices, and upstream contributor list are retained. Please also acknowledge RSL-RL: A Learning Library for Robotics Research when using this training infrastructure. Third-party license notices are in licenses/.

About

Diffusion-augmented MDPs in PyTorch: DA:PPO, DA:REPPO, DA:WPO and reproducible multimodal robotics demos

Topics

Resources

Stars

0 stars

Watchers

0 watching

Forks

Releases

Packages

Contributors

Languages