diff --git a/python/mujoco/rollout.cc b/python/mujoco/rollout.cc index e591af34..2bb4dbcf 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; diff --git a/python/mujoco/rollout_test.py b/python/mujoco/rollout_test.py index f27dc02b..f2b5358f 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 @@ -25,7 +26,6 @@ import numpy as np import mujoco from mujoco import rollout - # -------------------------- models used for testing --------------------------- TEST_XML = r""" @@ -875,6 +875,26 @@ 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 ----------------