Fix rollout notebook, update light.directional to light.type.

PiperOrigin-RevId: 773646227
Change-Id: I3bb54e548b1a74d7523aaa91be0fc3cb980866f3
This commit is contained in:
Tom Power
2025-06-20 04:43:06 -07:00
committed by Copybara-Service
parent 0488d9f4eb
commit 60955d7e88
+197 -78
View File
@@ -5,6 +5,7 @@
"id": "6adc68e0-a943-44ab-9af5-4bc62cc19f34",
"metadata": {
"editable": true,
"id": "6adc68e0-a943-44ab-9af5-4bc62cc19f34",
"tags": []
},
"source": [
@@ -16,7 +17,7 @@
"\n",
"This notebook describes the `rollout` module included in the MuJoCo Python library. It performs simulation \"rollouts\" with an underlying C++ function. The rollouts can be multithreaded.\n",
"\n",
"Below, the usage of each argument is explained with examples. Then some examples for advanced use cases are provided. Finally, `rollout` is benchmarked against pure python and MJX. \n",
"Below, the usage of each argument is explained with examples. Then some examples for advanced use cases are provided. Finally, `rollout` is benchmarked against pure python and MJX.\n",
"\n",
"Note the benchmarks were designed to run on >16 thread CPU and an RTX 4090 or A100. They do not run in a reasonable amount of time on a typical free colab runtime.\n",
"\n",
@@ -41,6 +42,7 @@
"id": "5d8a6604-0948-4a42-a48d-249c7f0c462b",
"metadata": {
"editable": true,
"id": "5d8a6604-0948-4a42-a48d-249c7f0c462b",
"tags": []
},
"source": [
@@ -53,6 +55,7 @@
"id": "0f9fbad1-59d0-40ac-b2b6-99f37313670f",
"metadata": {
"editable": true,
"id": "0f9fbad1-59d0-40ac-b2b6-99f37313670f",
"tags": [
"hide-input"
]
@@ -154,7 +157,9 @@
{
"cell_type": "markdown",
"id": "fc69d0f4",
"metadata": {},
"metadata": {
"id": "fc69d0f4"
},
"source": [
"# Helper Functions"
]
@@ -165,6 +170,7 @@
"id": "082482c7",
"metadata": {
"editable": true,
"id": "082482c7",
"tags": [
"hide-input"
]
@@ -251,7 +257,7 @@
" light.cutoff = 45\n",
" light.diffuse = [0.8, 0.8, 0.8]\n",
" light.dir = [0, 0, -1]\n",
" light.directional = 0\n",
" light.type = mujoco.mjtLightType.mjLIGHT_SPOT\n",
" light.exponent = 10\n",
" light.headlight = 0\n",
" light.specular = [0.3, 0.3, 0.3]\n",
@@ -267,7 +273,9 @@
{
"cell_type": "markdown",
"id": "c0570c2c",
"metadata": {},
"metadata": {
"id": "c0570c2c"
},
"source": [
"# Using `rollout`\n",
"\n",
@@ -280,9 +288,11 @@
},
{
"cell_type": "code",
"execution_count": null,
"execution_count": 0,
"id": "849b93e5",
"metadata": {},
"metadata": {
"id": "849b93e5"
},
"outputs": [],
"source": [
"#@title Benchmarked models\n",
@@ -368,7 +378,9 @@
{
"cell_type": "markdown",
"id": "55d171f7-541b-4441-aa18-da86d6716410",
"metadata": {},
"metadata": {
"id": "55d171f7-541b-4441-aa18-da86d6716410"
},
"source": [
"## Usage\n",
"\n",
@@ -380,9 +392,11 @@
},
{
"cell_type": "code",
"execution_count": null,
"execution_count": 0,
"id": "9cd2f94a-11df-4247-986c-5a56af69a1f5",
"metadata": {},
"metadata": {
"id": "9cd2f94a-11df-4247-986c-5a56af69a1f5"
},
"outputs": [],
"source": [
"print(rollout.rollout.__doc__)"
@@ -392,6 +406,7 @@
"cell_type": "markdown",
"id": "b6f7a094-8352-4b07-99ee-5278e3036cd5",
"metadata": {
"id": "b6f7a094-8352-4b07-99ee-5278e3036cd5",
"tags": []
},
"source": [
@@ -403,9 +418,11 @@
},
{
"cell_type": "code",
"execution_count": null,
"execution_count": 0,
"id": "849af5f2-9de1-4cb9-bc3a-c9b7acf0e3fe",
"metadata": {},
"metadata": {
"id": "849af5f2-9de1-4cb9-bc3a-c9b7acf0e3fe"
},
"outputs": [],
"source": [
"nbatch = 100 # Simulate this many tops\n",
@@ -437,16 +454,20 @@
{
"cell_type": "markdown",
"id": "aa2cf151-bf9a-4a23-b7fe-6a766979d93f",
"metadata": {},
"metadata": {
"id": "aa2cf151-bf9a-4a23-b7fe-6a766979d93f"
},
"source": [
"Our model has an angular velocity sensor the middle of the top. Let's plot the response using the `sensordata` array that rollout returns."
]
},
{
"cell_type": "code",
"execution_count": null,
"execution_count": 0,
"id": "957b8566-da31-410b-b385-e78241c5247a",
"metadata": {},
"metadata": {
"id": "957b8566-da31-410b-b385-e78241c5247a"
},
"outputs": [],
"source": [
"plt.figure(figsize=(12, 8))\n",
@@ -462,7 +483,9 @@
{
"cell_type": "markdown",
"id": "58044bc1-f98c-4bbf-a703-40ba075552a0",
"metadata": {},
"metadata": {
"id": "58044bc1-f98c-4bbf-a703-40ba075552a0"
},
"source": [
"### Example: different models\n",
"100 gray tops is kind of boring. It would be better if they were colorful and different sizes!\n",
@@ -474,9 +497,11 @@
},
{
"cell_type": "code",
"execution_count": null,
"execution_count": 0,
"id": "7c39e79e-8942-4fea-b306-ea0cb3c826e2",
"metadata": {},
"metadata": {
"id": "7c39e79e-8942-4fea-b306-ea0cb3c826e2"
},
"outputs": [],
"source": [
"# Make 100 tippe tops with different colors and sizes\n",
@@ -544,16 +569,20 @@
{
"cell_type": "markdown",
"id": "cf485c08-72be-4169-89b6-9d93df8ebbe3",
"metadata": {},
"metadata": {
"id": "cf485c08-72be-4169-89b6-9d93df8ebbe3"
},
"source": [
"Because the models are now different, the measurements of the gyro sensor are not consistent even though the initial state for each rollout was the same."
]
},
{
"cell_type": "code",
"execution_count": null,
"execution_count": 0,
"id": "b8a5d3d4-24e7-41a1-b3bd-7b63c1812b03",
"metadata": {},
"metadata": {
"id": "b8a5d3d4-24e7-41a1-b3bd-7b63c1812b03"
},
"outputs": [],
"source": [
"plt.figure(figsize=(12, 8))\n",
@@ -569,7 +598,9 @@
{
"cell_type": "markdown",
"id": "3841a669-6cd1-427e-a629-20a10a6e3a34",
"metadata": {},
"metadata": {
"id": "3841a669-6cd1-427e-a629-20a10a6e3a34"
},
"source": [
"### Example: control inputs\n",
"Open loop controls can be passed to `rollout` via the `control` argument. If passed, `nstep` no longer needs to be specified as it can be inferred from the size of `control`.\n",
@@ -579,9 +610,11 @@
},
{
"cell_type": "code",
"execution_count": null,
"execution_count": 0,
"id": "2a184873-8d24-45da-b444-8d21f5dcd733",
"metadata": {},
"metadata": {
"id": "2a184873-8d24-45da-b444-8d21f5dcd733"
},
"outputs": [],
"source": [
"# Episode parameters.\n",
@@ -632,7 +665,9 @@
{
"cell_type": "markdown",
"id": "4d89e4fb-7711-4d23-8ff3-eb5030fa8bf7",
"metadata": {},
"metadata": {
"id": "4d89e4fb-7711-4d23-8ff3-eb5030fa8bf7"
},
"source": [
"`rollout`'s `control_spec` argument can be used to indicate `control` contains values for actuators, generalized forces, cartesian forces, mocap poses, and/or the activation/deactivation of equality constraints. Internally, this is managed through [mj_setState](https://mujoco.readthedocs.io/en/stable/APIreference/APIfunctions.html#mj-setstate) and `control_spec` corresponds to `mj_setState`'s `spec` argument.\n",
"\n",
@@ -641,9 +676,11 @@
},
{
"cell_type": "code",
"execution_count": null,
"execution_count": 0,
"id": "4b02bb61-912d-47de-a956-aadfcd4c5cd5",
"metadata": {},
"metadata": {
"id": "4b02bb61-912d-47de-a956-aadfcd4c5cd5"
},
"outputs": [],
"source": [
"xfrc_size = mujoco.mj_stateSize(humanoid_model, mujoco.mjtState.mjSTATE_XFRC_APPLIED)\n",
@@ -676,7 +713,9 @@
{
"cell_type": "markdown",
"id": "0961c3ec-a691-4875-9a55-227a3d29c472",
"metadata": {},
"metadata": {
"id": "0961c3ec-a691-4875-9a55-227a3d29c472"
},
"source": [
"# Advanced usage"
]
@@ -684,7 +723,9 @@
{
"cell_type": "markdown",
"id": "VfYIyXWcLKfg",
"metadata": {},
"metadata": {
"id": "VfYIyXWcLKfg"
},
"source": [
"## skip_checks\n",
"\n",
@@ -706,9 +747,11 @@
},
{
"cell_type": "code",
"execution_count": null,
"execution_count": 0,
"id": "d02cc8e8-63cd-4852-ab3c-364a18025a95",
"metadata": {},
"metadata": {
"id": "d02cc8e8-63cd-4852-ab3c-364a18025a95"
},
"outputs": [],
"source": [
"nbatch = 1000\n",
@@ -752,7 +795,9 @@
{
"cell_type": "markdown",
"id": "92627030-4726-4689-be8b-f1ba75905104",
"metadata": {},
"metadata": {
"id": "92627030-4726-4689-be8b-f1ba75905104"
},
"source": [
"As expected, as `nstep` increases, the benefits of using skip checks fades quickly. However, at low nstep and high batch sizes, it can make a significant difference.\n",
"\n",
@@ -762,7 +807,9 @@
{
"cell_type": "markdown",
"id": "d32a77b5-24bd-4d17-80ac-15cc4d03731c",
"metadata": {},
"metadata": {
"id": "d32a77b5-24bd-4d17-80ac-15cc4d03731c"
},
"source": [
"## Reusing threadpools (`Rollout` class)\n",
"\n",
@@ -773,9 +820,11 @@
},
{
"cell_type": "code",
"execution_count": null,
"execution_count": 0,
"id": "dd05bbdf-f389-4e4e-b389-d47fe976cb49",
"metadata": {},
"metadata": {
"id": "dd05bbdf-f389-4e4e-b389-d47fe976cb49"
},
"outputs": [],
"source": [
"nbatch = 100\n",
@@ -813,7 +862,9 @@
{
"cell_type": "markdown",
"id": "9b3e14a1-71f3-430d-a3c2-1aadcf6c2671",
"metadata": {},
"metadata": {
"id": "9b3e14a1-71f3-430d-a3c2-1aadcf6c2671"
},
"source": [
"## Reusing threadpools (`rollout` method)\n",
"\n",
@@ -824,9 +875,11 @@
},
{
"cell_type": "code",
"execution_count": null,
"execution_count": 0,
"id": "b6aa6801",
"metadata": {},
"metadata": {
"id": "b6aa6801"
},
"outputs": [],
"source": [
"nbatch = 1000\n",
@@ -846,16 +899,20 @@
{
"cell_type": "markdown",
"id": "144378d3",
"metadata": {},
"metadata": {
"id": "144378d3"
},
"source": [
"Second, if `rollout` reuses the same threadpool between calls, it is no longer safe to call `rollout` from multiple threads. For example the following is not allowed (the offending lines are commented out to avoid crashing the interpreter):"
]
},
{
"cell_type": "code",
"execution_count": null,
"execution_count": 0,
"id": "7f46a6d8",
"metadata": {},
"metadata": {
"id": "7f46a6d8"
},
"outputs": [],
"source": [
"thread1 = threading.Thread(target=lambda: rollout.rollout(top_model, top_datas, initial_states, nstep=nstep, persistent_pool=True))\n",
@@ -871,7 +928,9 @@
{
"cell_type": "markdown",
"id": "78c1f864-5238-4e27-a7ec-d03c45484d9a",
"metadata": {},
"metadata": {
"id": "78c1f864-5238-4e27-a7ec-d03c45484d9a"
},
"source": [
"## chunk_size\n",
"\n",
@@ -882,9 +941,11 @@
},
{
"cell_type": "code",
"execution_count": null,
"execution_count": 0,
"id": "a1be8f93",
"metadata": {},
"metadata": {
"id": "a1be8f93"
},
"outputs": [],
"source": [
"nbatch = 100\n",
@@ -928,7 +989,9 @@
{
"cell_type": "markdown",
"id": "84c746e0-5d2e-47dc-ac76-f2b6f790b7c7",
"metadata": {},
"metadata": {
"id": "84c746e0-5d2e-47dc-ac76-f2b6f790b7c7"
},
"source": [
"## Warmstarting\n",
"\n",
@@ -941,9 +1004,11 @@
},
{
"cell_type": "code",
"execution_count": null,
"execution_count": 0,
"id": "d4d9f660-f83c-432e-a579-124a7ecab4fb",
"metadata": {},
"metadata": {
"id": "d4d9f660-f83c-432e-a579-124a7ecab4fb"
},
"outputs": [],
"source": [
"top_model_cg = copy.copy(top_model)\n",
@@ -1004,7 +1069,9 @@
{
"cell_type": "markdown",
"id": "7c2cf4fa",
"metadata": {},
"metadata": {
"id": "7c2cf4fa"
},
"source": [
"As expected, the middle animation (with warmstarting) matches the continuous rollout on the left. However, the model that did not use warmstarting diverged."
]
@@ -1012,7 +1079,9 @@
{
"cell_type": "markdown",
"id": "7944637f",
"metadata": {},
"metadata": {
"id": "7944637f"
},
"source": [
"# Benchmarks\n",
"\n",
@@ -1029,9 +1098,11 @@
},
{
"cell_type": "code",
"execution_count": null,
"execution_count": 0,
"id": "cb6355dd",
"metadata": {},
"metadata": {
"id": "cb6355dd"
},
"outputs": [],
"source": [
"def python_rollout(model, data, nbatch, nstep):\n",
@@ -1043,7 +1114,9 @@
{
"cell_type": "markdown",
"id": "6fe4a78b",
"metadata": {},
"metadata": {
"id": "6fe4a78b"
},
"source": [
"To run nbatch rollouts with `rollout`, we need to make an array of nbatch initial states to start the rollouts from.\n",
"\n",
@@ -1055,9 +1128,11 @@
},
{
"cell_type": "code",
"execution_count": null,
"execution_count": 0,
"id": "74f143e2",
"metadata": {},
"metadata": {
"id": "74f143e2"
},
"outputs": [],
"source": [
"def nthread_rollout(model, data, nbatch, nstep, nthread, rollout_):\n",
@@ -1071,7 +1146,9 @@
{
"cell_type": "markdown",
"id": "b75dc44c",
"metadata": {},
"metadata": {
"id": "b75dc44c"
},
"source": [
"Next, we benchmark the Python loop and `rollout` in both single threaded and multithreaded modes. The three benchmarks take about 2.5 minutes in total to run in total on an AMD 5800X3D."
]
@@ -1081,7 +1158,8 @@
"execution_count": 0,
"id": "0301e3ee",
"metadata": {
"cellView": "form"
"cellView": "form",
"id": "0301e3ee"
},
"outputs": [],
"source": [
@@ -1204,16 +1282,20 @@
{
"cell_type": "markdown",
"id": "08fb0c12",
"metadata": {},
"metadata": {
"id": "08fb0c12"
},
"source": [
"### Tippe Top Benchmark"
]
},
{
"cell_type": "code",
"execution_count": null,
"execution_count": 0,
"id": "f7e54830",
"metadata": {},
"metadata": {
"id": "f7e54830"
},
"outputs": [],
"source": [
"nominal_nbatch = 256 # Batch size to use when testing different nstep\n",
@@ -1232,16 +1314,20 @@
{
"cell_type": "markdown",
"id": "edefb26e",
"metadata": {},
"metadata": {
"id": "edefb26e"
},
"source": [
"### Humanoid Benchmark"
]
},
{
"cell_type": "code",
"execution_count": null,
"execution_count": 0,
"id": "c9e58c6c",
"metadata": {},
"metadata": {
"id": "c9e58c6c"
},
"outputs": [],
"source": [
"nominal_nbatch = 256 # Batch size to use when testing different nstep\n",
@@ -1260,16 +1346,20 @@
{
"cell_type": "markdown",
"id": "468903bb",
"metadata": {},
"metadata": {
"id": "468903bb"
},
"source": [
"### Humanoid100 Benchmark"
]
},
{
"cell_type": "code",
"execution_count": null,
"execution_count": 0,
"id": "83d775d4",
"metadata": {},
"metadata": {
"id": "83d775d4"
},
"outputs": [],
"source": [
"nominal_nbatch = 128 # Batch size to use when testing different nstep\n",
@@ -1293,7 +1383,9 @@
{
"cell_type": "markdown",
"id": "d1133084",
"metadata": {},
"metadata": {
"id": "d1133084"
},
"source": [
"# MJX versus `rollout`"
]
@@ -1301,7 +1393,9 @@
{
"cell_type": "markdown",
"id": "c1638f2d",
"metadata": {},
"metadata": {
"id": "c1638f2d"
},
"source": [
"Next we will benchmark `rollout` and MJX using the tippe top and humanoid models (humanoid100 is not supported by MJX).\n",
"\n",
@@ -1315,7 +1409,8 @@
"execution_count": 0,
"id": "7c86d157",
"metadata": {
"cellView": "form"
"cellView": "form",
"id": "7c86d157"
},
"outputs": [],
"source": [
@@ -1436,16 +1531,20 @@
{
"cell_type": "markdown",
"id": "a2dafd2e",
"metadata": {},
"metadata": {
"id": "a2dafd2e"
},
"source": [
"### MJX Tippe Top Benchmark"
]
},
{
"cell_type": "code",
"execution_count": null,
"execution_count": 0,
"id": "98c580b0",
"metadata": {},
"metadata": {
"id": "98c580b0"
},
"outputs": [],
"source": [
"nominal_nbatch = 16384 # Batch size to use when testing different nstep\n",
@@ -1461,16 +1560,20 @@
{
"cell_type": "markdown",
"id": "205da5da",
"metadata": {},
"metadata": {
"id": "205da5da"
},
"source": [
"### MJX Humanoid Benchmark"
]
},
{
"cell_type": "code",
"execution_count": null,
"execution_count": 0,
"id": "53166ae1",
"metadata": {},
"metadata": {
"id": "53166ae1"
},
"outputs": [],
"source": [
"nominal_nbatch = 4096 # Batch size to use when testing different nstep\n",
@@ -1486,7 +1589,9 @@
{
"cell_type": "markdown",
"id": "fb2caa72",
"metadata": {},
"metadata": {
"id": "fb2caa72"
},
"source": [
"### MJX Multiple Humanoids in one model\n",
"\n",
@@ -1499,9 +1604,11 @@
},
{
"cell_type": "code",
"execution_count": null,
"execution_count": 0,
"id": "3d6be608",
"metadata": {},
"metadata": {
"id": "3d6be608"
},
"outputs": [],
"source": [
"max_humanoids = 10\n",
@@ -1549,9 +1656,11 @@
},
{
"cell_type": "code",
"execution_count": null,
"execution_count": 0,
"id": "b6c5fc2e",
"metadata": {},
"metadata": {
"id": "b6c5fc2e"
},
"outputs": [],
"source": [
"#@title Plot MJX nhumanoid benchmark\n",
@@ -1584,6 +1693,16 @@
"\n",
"plot_mjx_nhumanoid_benchmark(t_rollout, t_mjx, nbatch, nstep, max_humanoids)"
]
},
{
"cell_type": "code",
"execution_count": 0,
"id": "UW0aoKXK7ALd",
"metadata": {
"id": "UW0aoKXK7ALd"
},
"outputs": [],
"source": []
}
],
"metadata": {