diff --git a/mjx/tutorial.ipynb b/mjx/tutorial.ipynb index 65579cad..f57aee4d 100644 --- a/mjx/tutorial.ipynb +++ b/mjx/tutorial.ipynb @@ -367,7 +367,7 @@ "rng = jax.random.split(rng, 4096)\n", "batch = jax.vmap(lambda rng: mjx_data.replace(qpos=jax.random.uniform(rng, (1,))))(rng)\n", "\n", - "jit_step = jax.vmap(mjx.step, in_axes=(None, 0))\n", + "jit_step = jax.jit(jax.vmap(mjx.step, in_axes=(None, 0)))\n", "batch = jit_step(mjx_model, batch)\n", "\n", "print(batch.qpos)"