rollout fix cosmetics

This commit is contained in:
Levi Burner
2025-02-21 08:44:17 -05:00
parent aca4e5d6ae
commit a0971ef749
2 changed files with 9 additions and 5 deletions
+1 -1
View File
@@ -227,7 +227,7 @@ mjtNum* get_array_ptr(std::optional<const py::array_t<mjtNum>> arg,
py::buffer_info info = arg->request();
// check size
size_t expected_size =
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;
+8 -4
View File
@@ -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 ----------------