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 ----------------