Merge pull request #2439 from aftersomemath:rollout-size_t

PiperOrigin-RevId: 729519951
Change-Id: Ic50b08960da84d2277820212d19fcbab9f2d67cf
This commit is contained in:
Copybara-Service
2025-02-21 07:17:25 -08:00
2 changed files with 32 additions and 10 deletions
+11 -9
View File
@@ -75,10 +75,11 @@ void _unsafe_rollout(std::vector<const mjModel*>& 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<size_t>(mj_stateSize(m[0], mjSTATE_FULLPHYSICS));
size_t ncontrol = static_cast<size_t>(mj_stateSize(m[0], control_spec));
size_t nv = static_cast<size_t>(m[0]->nv);
int nbody = m[0]->nbody, neq = m[0]->neq;
size_t nsensordata = static_cast<size_t>(m[0]->nsensordata);
// clear user inputs if unspecified
if (!(control_spec & mjSTATE_CTRL)) {
@@ -92,7 +93,7 @@ void _unsafe_rollout(std::vector<const mjModel*>& 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<const mjModel*>& 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<const mjModel*>& 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<size_t>(nstep) + t;
if (state) {
mj_getState(m[r], d, state + step*nstate, mjSTATE_FULLPHYSICS);
}
@@ -152,7 +153,7 @@ void _unsafe_rollout(std::vector<const mjModel*>& m, mjData* d, int start_roll,
break;
}
int step = r*nstep + t;
size_t step = r*static_cast<size_t>(nstep) + t;
// controls
if (control) {
@@ -226,7 +227,8 @@ mjtNum* get_array_ptr(std::optional<const py::array_t<mjtNum>> arg,
py::buffer_info info = arg->request();
// check size
int expected_size = nbatch * nstep * dim;
size_t expected_size =
static_cast<size_t>(nbatch) * static_cast<size_t>(nstep) * static_cast<size_t>(dim);
if (info.size != expected_size) {
std::ostringstream msg;
msg << name << ".size should be " << expected_size << ", got " << info.size;
+21 -1
View File
@@ -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 ----------------