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.

Downloads last month

-

Downloads are not tracked for this model. How to track
Video Preview
loading