From df71467eedc10a1f20e1caf3e36b8236dbd5bee3 Mon Sep 17 00:00:00 2001 From: Levi Burner Date: Mon, 3 Feb 2025 22:22:37 -0500 Subject: [PATCH 1/3] rollout use size_t for array size check and pointer arithmetic --- python/mujoco/rollout.cc | 20 +++++++++++--------- 1 file changed, 11 insertions(+), 9 deletions(-) diff --git a/python/mujoco/rollout.cc b/python/mujoco/rollout.cc index e591af34..08e24450 100644 --- a/python/mujoco/rollout.cc +++ b/python/mujoco/rollout.cc @@ -75,10 +75,11 @@ void _unsafe_rollout(std::vector& m, mjData* d, int start_roll, const mjtNum* state0, const mjtNum* warmstart0, const mjtNum* control, mjtNum* state, mjtNum* sensordata) { // sizes - int nstate = mj_stateSize(m[0], mjSTATE_FULLPHYSICS); - int ncontrol = mj_stateSize(m[0], control_spec); - int nv = m[0]->nv, nbody = m[0]->nbody, neq = m[0]->neq; - int nsensordata = m[0]->nsensordata; + size_t nstate = static_cast(mj_stateSize(m[0], mjSTATE_FULLPHYSICS)); + size_t ncontrol = static_cast(mj_stateSize(m[0], control_spec)); + size_t nv = static_cast(m[0]->nv); + int nbody = m[0]->nbody, neq = m[0]->neq; + size_t nsensordata = static_cast(m[0]->nsensordata); // clear user inputs if unspecified if (!(control_spec & mjSTATE_CTRL)) { @@ -92,7 +93,7 @@ void _unsafe_rollout(std::vector& m, mjData* d, int start_roll, } // loop over rollouts - for (int r = start_roll; r < end_roll; r++) { + for (size_t r = start_roll; r < end_roll; r++) { // clear user inputs if unspecified if (!(control_spec & mjSTATE_MOCAP_POS)) { for (int i = 0; i < nbody; i++) { @@ -128,7 +129,7 @@ void _unsafe_rollout(std::vector& m, mjData* d, int start_roll, } // roll out trajectory - for (int t = 0; t < nstep; t++) { + for (size_t t = 0; t < nstep; t++) { // check for warnings bool nwarning = false; for (int i = 0; i < mjNWARNING; i++) { @@ -141,7 +142,7 @@ void _unsafe_rollout(std::vector& m, mjData* d, int start_roll, // if any warnings, fill remaining outputs with current outputs, break if (nwarning) { for (; t < nstep; t++) { - int step = r*nstep + t; + size_t step = r*static_cast(nstep) + t; if (state) { mj_getState(m[r], d, state + step*nstate, mjSTATE_FULLPHYSICS); } @@ -152,7 +153,7 @@ void _unsafe_rollout(std::vector& m, mjData* d, int start_roll, break; } - int step = r*nstep + t; + size_t step = r*static_cast(nstep) + t; // controls if (control) { @@ -226,7 +227,8 @@ mjtNum* get_array_ptr(std::optional> arg, py::buffer_info info = arg->request(); // check size - int expected_size = nbatch * nstep * dim; + size_t expected_size = + static_cast(nbatch) * static_cast(nstep) * static_cast(dim); if (info.size != expected_size) { std::ostringstream msg; msg << name << ".size should be " << expected_size << ", got " << info.size; From aca4e5d6aef8c8c89cfe6a17d4289edc49dbc667 Mon Sep 17 00:00:00 2001 From: Levi Burner Date: Thu, 20 Feb 2025 21:12:52 -0500 Subject: [PATCH 2/3] rollout test for integer overflow in array size check --- python/mujoco/rollout_test.py | 16 ++++++++++++++++ 1 file changed, 16 insertions(+) diff --git a/python/mujoco/rollout_test.py b/python/mujoco/rollout_test.py index f27dc02b..f37c2244 100644 --- a/python/mujoco/rollout_test.py +++ b/python/mujoco/rollout_test.py @@ -16,6 +16,7 @@ import concurrent.futures import copy +import os import threading from absl.testing import absltest @@ -875,6 +876,21 @@ class MuJoCoRolloutTest(parameterized.TestCase): model, [copy.copy(data) for i in range(3)], initial_state, control ) + @absltest.skip(reason='Takes a long time to run') + def test_large_state(self): + model = mujoco.MjModel.from_xml_string(TEST_XML) + nstate = mujoco.mj_stateSize(model, mujoco.mjtState.mjSTATE_FULLPHYSICS) + data = mujoco.MjData(model) + + nthread = os.cpu_count() + nbatch = nthread + + nstep = ((2**31) // (nstate*nbatch)) + 2 + assert nstep * nstate * nbatch > 2**31 + + initial_state = np.random.randn(nbatch, nstate) + rollout.rollout(model, [copy.copy(data) for _ in range(nthread)], + initial_state, nstep=nstep) # -------------- Python implementation of rollout functionality ---------------- From a0971ef749ead0668a072d9cc803f23fbc1ee7d7 Mon Sep 17 00:00:00 2001 From: Levi Burner Date: Fri, 21 Feb 2025 08:44:17 -0500 Subject: [PATCH 3/3] rollout fix cosmetics --- python/mujoco/rollout.cc | 2 +- python/mujoco/rollout_test.py | 12 ++++++++---- 2 files changed, 9 insertions(+), 5 deletions(-) diff --git a/python/mujoco/rollout.cc b/python/mujoco/rollout.cc index 08e24450..2bb4dbcf 100644 --- a/python/mujoco/rollout.cc +++ b/python/mujoco/rollout.cc @@ -227,7 +227,7 @@ mjtNum* get_array_ptr(std::optional> arg, py::buffer_info info = arg->request(); // check size - size_t expected_size = + size_t expected_size = static_cast(nbatch) * static_cast(nstep) * static_cast(dim); if (info.size != expected_size) { std::ostringstream msg; diff --git a/python/mujoco/rollout_test.py b/python/mujoco/rollout_test.py index f37c2244..f2b5358f 100644 --- a/python/mujoco/rollout_test.py +++ b/python/mujoco/rollout_test.py @@ -26,7 +26,6 @@ import numpy as np import mujoco from mujoco import rollout - # -------------------------- models used for testing --------------------------- TEST_XML = r""" @@ -885,12 +884,17 @@ class MuJoCoRolloutTest(parameterized.TestCase): nthread = os.cpu_count() nbatch = nthread - nstep = ((2**31) // (nstate*nbatch)) + 2 + nstep = ((2**31) // (nstate * nbatch)) + 2 assert nstep * nstate * nbatch > 2**31 initial_state = np.random.randn(nbatch, nstate) - rollout.rollout(model, [copy.copy(data) for _ in range(nthread)], - initial_state, nstep=nstep) + rollout.rollout( + model, + [copy.copy(data) for _ in range(nthread)], + initial_state, + nstep=nstep, + ) + # -------------- Python implementation of rollout functionality ----------------