MIMIC-MJX checkpoints
Trained imitation policies for track-mjx, part of MIMIC-MJX. Reference motion lives in the companion dataset, talmolab/MIMIC-MJX.
track-mjx v1.1 checkpoints
<body>/track-mjx-v1.1/<backend>/ holds one policy per body model and MuJoCo
backend (jax = MJX, warp = MuJoCo Warp). Each was trained with the bundled
track-mjx config for that body; the only override was the backend
(env_config.mujoco_impl), plus train_setup.train_config.num_envs=2048 for the
worm on Warp. The full training config is stored inside each checkpoint.
| Body | Folder | Config | Steps | Final eval reward, Warp / MJX (train clips) |
|---|---|---|---|---|
| Rodent | rodent/ |
rodent-full-clips |
3.0 B | 1783 / 1780 |
| Fruit fly | fruitfly/ |
fly |
2.0 B | 2567 / 2580 |
| Mouse arm | arm/ |
mouse-arm |
650 M | 2950 / 2950 |
| Stick insect | stick/ |
stick |
2.0 B | 2048 / 2048 |
| Worm | worm/ |
celegans |
1.5 B | 433 / 435 |
In evaluation the policies track the full episode with at most 0.8% early terminations, and held-out clip rewards are within 1.5% of training-clip rewards. The mouse-arm MJX checkpoint is step 24 of 25 (625 M steps); the final save did not complete.
Requirements: track-mjx with vnl-playground at
79da61d or later
(the stick policies use the mesh-based Sungaya walker), and the reference data
downloaded into the track-mjx data/ directory. The worm uses
data/worm/celegans_ik_only_04182019am_centerline_locomotion_2d_scale-0.1.h5.
hf download talmolab/MIMIC-MJX --include 'rodent/track-mjx-v1.1/warp/**' --local-dir checkpoints
from track_mjx.agent import checkpointing
from track_mjx.config import utils as config_utils
ckpt = checkpointing.load_checkpoint_for_eval("checkpoints/rodent/track-mjx-v1.1/warp")
cfg, _, _ = config_utils.prepare_config(ckpt["cfg"])
inference_fn = checkpointing.load_inference_fn(cfg, ckpt["policy"])
See notebooks/rollout_from_checkpoint.ipynb in track-mjx for a full rollout.
The other folders hold earlier checkpoints from the MIMIC-MJX paper.