Support float32 in mujoco.rollout Python wrapper and tests.

- Update rollout.py to use mujoco.MJTNUM_DTYPE instead of np.float64.
- Update rollout_test.py to use DTYPE for all state/control/sensor arrays to avoid copies and type mismatches.

PiperOrigin-RevId: 928539946
Change-Id: Iee2250c049de799cb0e23ffe5558fe6eea8b3fd8
This commit is contained in:
Yuval Tassa
2026-06-08 06:58:54 -07:00
committed by Copybara-Service
parent 116cfe9c70
commit a0f014593b
3 changed files with 114 additions and 96 deletions
+9 -7
View File
@@ -13,6 +13,7 @@
// limitations under the License.
#include <algorithm>
#include <cstddef>
#include <iostream>
#include <memory>
#include <optional>
@@ -177,10 +178,11 @@ void _unsafe_rollout(std::vector<const mjModel*>& m, mjData* d, int start_roll,
}
// C-style threaded version of _unsafe_rollout
void _unsafe_rollout_threaded(std::vector<const mjModel*>& m, std::vector<mjData*>& d,
int nbatch, int nstep, unsigned int control_spec,
const mjtNum* state0, const mjtNum* warmstart0,
const mjtNum* control, mjtNum* state, mjtNum* sensordata,
void _unsafe_rollout_threaded(std::vector<const mjModel*>& m,
std::vector<mjData*>& d, int nbatch, int nstep,
unsigned int control_spec, const mjtNum* state0,
const mjtNum* warmstart0, const mjtNum* control,
mjtNum* state, mjtNum* sensordata,
ThreadPool* pool, int chunk_size) {
int nfulljobs = nbatch / chunk_size;
int chunk_remainder = nbatch % chunk_size;
@@ -209,7 +211,7 @@ void _unsafe_rollout_threaded(std::vector<const mjModel*>& m, std::vector<mjData
pool->Schedule(task);
}
// wait for job counter to incremented up to the number of jobs submitted by this thread
// wait for counter to increment up to number of jobs submitted by this thread
pool->WaitCount(njobs);
}
@@ -227,8 +229,8 @@ mjtNum* get_array_ptr(std::optional<const py::array_t<mjtNum>> arg,
py::buffer_info info = arg->request();
// check size
size_t expected_size =
static_cast<size_t>(nbatch) * static_cast<size_t>(nstep) * static_cast<size_t>(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;
+4 -4
View File
@@ -219,9 +219,9 @@ class Rollout:
# allocate output if not provided
if state is None:
state = np.empty((nbatch, nstep, nstate))
state = np.empty((nbatch, nstep, nstate), dtype=mujoco.MJTNUM_DTYPE)
if sensordata is None:
sensordata = np.empty((nbatch, nstep, nsensordata))
sensordata = np.empty((nbatch, nstep, nsensordata), dtype=mujoco.MJTNUM_DTYPE)
# call rollout
self.rollout_.rollout(
@@ -376,7 +376,7 @@ def _ensure_2d(arg):
if arg is None:
return None
else:
return np.ascontiguousarray(np.atleast_2d(arg), dtype=np.float64)
return np.ascontiguousarray(np.atleast_2d(arg), dtype=mujoco.MJTNUM_DTYPE)
def _ensure_3d(arg):
@@ -390,7 +390,7 @@ def _ensure_3d(arg):
arg = arg[np.newaxis, np.newaxis, ...]
elif arg.ndim == 2:
arg = arg[np.newaxis, ...]
return np.ascontiguousarray(arg, dtype=np.float64)
return np.ascontiguousarray(arg, dtype=mujoco.MJTNUM_DTYPE)
def _infer_dimension(dim, value, **kwargs):
+101 -85
View File
@@ -21,11 +21,27 @@ import threading
from absl.testing import absltest
from absl.testing import parameterized
import numpy as np
import mujoco
from mujoco import rollout
import numpy as np
DTYPE = mujoco.MJTNUM_DTYPE
def randn(*shape):
return np.random.standard_normal(shape).astype(DTYPE)
def empty(*shape):
return np.empty(shape, dtype=DTYPE)
def zeros(*shape):
return np.zeros(shape, dtype=DTYPE)
# -------------------------- models used for testing ---------------------------
TEST_XML = r"""
@@ -154,8 +170,8 @@ class MuJoCoRolloutTest(parameterized.TestCase):
nstate = mujoco.mj_stateSize(model, mujoco.mjtState.mjSTATE_FULLPHYSICS)
data = mujoco.MjData(model)
initial_state = np.random.randn(nstate)
control = np.random.randn(model.nu)
initial_state = randn(nstate)
control = randn(model.nu)
state, sensordata = rollout.rollout(model, data, initial_state, control)
mujoco.mj_resetData(model, data)
@@ -171,8 +187,8 @@ class MuJoCoRolloutTest(parameterized.TestCase):
nstate = mujoco.mj_stateSize(model, mujoco.mjtState.mjSTATE_FULLPHYSICS)
data = mujoco.MjData(model)
initial_state = np.random.randn(nstate)
control = np.random.randn(nstep, model.nu)
initial_state = randn(nstate)
control = randn(nstep, model.nu)
state, sensordata = rollout.rollout(model, data, initial_state, control)
py_state, py_sensordata = py_rollout(model, data, initial_state, control)
@@ -188,8 +204,8 @@ class MuJoCoRolloutTest(parameterized.TestCase):
nbatch = 5 # number of rollouts
nstep = 1 # number of steps
initial_state = np.random.randn(nbatch, nstate)
control = np.random.randn(nbatch, nstep, model.nu)
initial_state = randn(nbatch, nstate)
control = randn(nbatch, nstep, model.nu)
state, sensordata = rollout.rollout(model, data, initial_state, control)
mujoco.mj_resetData(model, data)
@@ -206,8 +222,8 @@ class MuJoCoRolloutTest(parameterized.TestCase):
nbatch = 5 # number of rollouts
nstep = 1 # number of steps
initial_state = np.random.randn(nbatch, nstate)
control = np.random.randn(nstep, model.nu)
initial_state = randn(nbatch, nstate)
control = randn(nstep, model.nu)
state, sensordata = rollout.rollout(model, data, initial_state, control)
mujoco.mj_resetData(model, data)
@@ -225,8 +241,8 @@ class MuJoCoRolloutTest(parameterized.TestCase):
nbatch = 5 # number of rollouts
nstep = 1 # number of steps
initial_state = np.random.randn(nstate)
control = np.random.randn(nbatch, nstep, model.nu)
initial_state = randn(nstate)
control = randn(nbatch, nstep, model.nu)
state, sensordata = rollout.rollout(model, data, initial_state, control)
mujoco.mj_resetData(model, data)
@@ -244,8 +260,8 @@ class MuJoCoRolloutTest(parameterized.TestCase):
nbatch = 5 # number of rollouts
nstep = 1 # number of steps
initial_state = np.random.randn(nstate)
control = np.random.randn(nstep, model.nu)
initial_state = randn(nstate)
control = randn(nstep, model.nu)
initial_warmstart = np.tile(data.qacc_warmstart.copy(), (nbatch, 1))
state, sensordata = rollout.rollout(
model, data, initial_state, control, initial_warmstart=initial_warmstart
@@ -267,9 +283,9 @@ class MuJoCoRolloutTest(parameterized.TestCase):
nbatch = 5 # number of rollouts
nstep = 1 # number of steps
initial_state = np.random.randn(nstate)
control = np.random.randn(nstep, model.nu)
state = np.empty((nbatch, nstep, nstate))
initial_state = randn(nstate)
control = randn(nstep, model.nu)
state = empty(nbatch, nstep, nstate)
state, sensordata = rollout.rollout(
model, data, initial_state, control, state=state
)
@@ -290,9 +306,9 @@ class MuJoCoRolloutTest(parameterized.TestCase):
nbatch = 5 # number of rollouts
nstep = 1 # number of steps
initial_state = np.random.randn(nstate)
control = np.random.randn(nstep, model.nu)
sensordata = np.empty((nbatch, nstep, model.nsensordata))
initial_state = randn(nstate)
control = randn(nstep, model.nu)
sensordata = empty(nbatch, nstep, model.nsensordata)
state, sensordata = rollout.rollout(
model, data, initial_state, control, sensordata=sensordata
)
@@ -313,10 +329,10 @@ class MuJoCoRolloutTest(parameterized.TestCase):
nbatch = 1 # number of rollouts
nstep = 3 # number of steps
initial_state = np.random.randn(nstate)
control = np.random.randn(model.nu)
state = np.empty((nbatch, nstep, nstate))
sensordata = np.empty((nbatch, nstep, model.nsensordata))
initial_state = randn(nstate)
control = randn(model.nu)
state = empty(nbatch, nstep, nstate)
sensordata = empty(nbatch, nstep, model.nsensordata)
rollout.rollout(
model, data, initial_state, control, state=state, sensordata=sensordata
)
@@ -335,8 +351,8 @@ class MuJoCoRolloutTest(parameterized.TestCase):
nbatch = 2 # number of initial states
nstep = 3 # number of timesteps
initial_state = np.random.randn(nbatch, nstate)
control = np.random.randn(nbatch, nstep, model.nu)
initial_state = randn(nbatch, nstate)
control = randn(nbatch, nstep, model.nu)
state, sensordata = rollout.rollout(model, data, initial_state, control)
py_state, py_sensordata = py_rollout(model, data, initial_state, control)
@@ -363,8 +379,8 @@ class MuJoCoRolloutTest(parameterized.TestCase):
nstate = mujoco.mj_stateSize(model[0], mujoco.mjtState.mjSTATE_FULLPHYSICS)
data = mujoco.MjData(model[0])
initial_state = np.random.randn(nbatch, nstate)
control = np.random.randn(nbatch, nstep, model[0].nu)
initial_state = randn(nbatch, nstate)
control = randn(nbatch, nstep, model[0].nu)
state, sensordata = rollout.rollout(model, data, initial_state, control)
py_state, py_sensordata = py_rollout(model, data, initial_state, control)
@@ -380,9 +396,9 @@ class MuJoCoRolloutTest(parameterized.TestCase):
nbatch = 2 # number of rollouts
nstep = 3 # number of timesteps
initial_state = np.random.randn(nbatch, nstate)
control = np.random.randn(nbatch, 1, model.nu)
state = np.empty((nbatch, nstep, nstate))
initial_state = randn(nbatch, nstate)
control = randn(nbatch, 1, model.nu)
state = empty(nbatch, nstep, nstate)
state, sensordata = rollout.rollout(
model, data, initial_state, control, state=state
)
@@ -401,7 +417,7 @@ class MuJoCoRolloutTest(parameterized.TestCase):
nbatch = 4 # number of rollouts
nstep = 3 # number of timesteps
initial_state = np.random.randn(nbatch, nstate)
initial_state = randn(nbatch, nstate)
control_spec = (
mujoco.mjtState.mjSTATE_CTRL
@@ -409,7 +425,7 @@ class MuJoCoRolloutTest(parameterized.TestCase):
| mujoco.mjtState.mjSTATE_XFRC_APPLIED
)
ncontrol = mujoco.mj_stateSize(model, control_spec)
control = np.random.randn(nbatch, nstep, ncontrol)
control = randn(nbatch, nstep, ncontrol)
state, sensordata = rollout.rollout(
model, data, initial_state, control, control_spec=control_spec
@@ -427,7 +443,7 @@ class MuJoCoRolloutTest(parameterized.TestCase):
data = mujoco.MjData(model)
nbatch = 4 # number of rollouts
initial_state = np.empty((nbatch, nstate))
initial_state = empty(nbatch, nstate)
# get diverging (0, 2) and non-diverging (1, 3) states
mujoco.mj_getState(
@@ -446,7 +462,7 @@ class MuJoCoRolloutTest(parameterized.TestCase):
nstep = 10000 # divergence after ~15s, timestep = 2e-3
state = np.random.randn(nbatch, nstep, nstate)
state = randn(nbatch, nstep, nstate)
rollout.rollout(model, data, initial_state, state=state)
@@ -466,10 +482,10 @@ class MuJoCoRolloutTest(parameterized.TestCase):
num_workers = 32
nbatch = 100
nstep = 5
initial_state = np.random.randn(nbatch, nstate)
state = np.empty((nbatch, nstep, nstate))
sensordata = np.empty((nbatch, nstep, model.nsensordata))
control = np.random.randn(nbatch, nstep, model.nu)
initial_state = randn(nbatch, nstate)
state = empty(nbatch, nstep, nstate)
sensordata = empty(nbatch, nstep, model.nsensordata)
control = randn(nbatch, nstep, model.nu)
thread_local = threading.local()
@@ -528,10 +544,10 @@ class MuJoCoRolloutTest(parameterized.TestCase):
num_workers = 32
nbatch = 100
nstep = 5
initial_state = np.random.randn(nbatch, nstate)
state = np.empty((nbatch, nstep, nstate))
sensordata = np.empty((nbatch, nstep, model.nsensordata))
control = np.random.randn(nbatch, nstep, model.nu)
initial_state = randn(nbatch, nstate)
state = empty(nbatch, nstep, nstate)
sensordata = empty(nbatch, nstep, model.nsensordata)
control = randn(nbatch, nstep, model.nu)
model_list = [copy.copy(model) for _ in range(nbatch)]
data_list = [mujoco.MjData(model) for _ in range(num_workers)]
@@ -557,10 +573,10 @@ class MuJoCoRolloutTest(parameterized.TestCase):
num_workers = 32
nbatch = 100
nstep = 5
initial_state = np.random.randn(nbatch, nstate)
state = np.empty((nbatch, nstep, nstate))
sensordata = np.empty((nbatch, nstep, model.nsensordata))
control = np.random.randn(nbatch, nstep, model.nu)
initial_state = randn(nbatch, nstate)
state = empty(nbatch, nstep, nstate)
sensordata = empty(nbatch, nstep, model.nsensordata)
control = randn(nbatch, nstep, model.nu)
model_list = [copy.copy(model) for _ in range(nbatch)]
data_list = [mujoco.MjData(model) for _ in range(num_workers)]
@@ -606,10 +622,10 @@ class MuJoCoRolloutTest(parameterized.TestCase):
num_workers = 32
nbatch = 100
nstep = 5
initial_state = np.random.randn(nbatch, nstate)
state = np.empty((nbatch, nstep, nstate))
sensordata = np.empty((nbatch, nstep, model.nsensordata))
control = np.random.randn(nbatch, nstep, model.nu)
initial_state = randn(nbatch, nstate)
state = empty(nbatch, nstep, nstate)
sensordata = empty(nbatch, nstep, model.nsensordata)
control = randn(nbatch, nstep, model.nu)
model_list = [copy.copy(model) for _ in range(nbatch)]
data_list = [mujoco.MjData(model) for _ in range(num_workers)]
@@ -640,8 +656,8 @@ class MuJoCoRolloutTest(parameterized.TestCase):
data = mujoco.MjData(model)
# take one step, save the state
state0 = np.zeros(nstate)
control = np.zeros(model.nu)
state0 = zeros(nstate)
control = zeros(model.nu)
state1, _ = step(model, data, state0, control)
# save qacc_warmstart
@@ -671,17 +687,17 @@ class MuJoCoRolloutTest(parameterized.TestCase):
nstate = mujoco.mj_stateSize(model, mujoco.mjtState.mjSTATE_FULLPHYSICS)
data = mujoco.MjData(model)
initial_state = np.zeros(nstate)
initial_state = zeros(nstate)
control_spec = (
mujoco.mjtState.mjSTATE_MOCAP_POS | mujoco.mjtState.mjSTATE_MOCAP_QUAT
)
pos1 = np.array((1.0, 2.0, 3.0))
quat1 = np.array((1.0, 2.0, 3.0, 4.0))
pos1 = np.array((1.0, 2.0, 3.0), dtype=DTYPE)
quat1 = np.array((1.0, 2.0, 3.0, 4.0), dtype=DTYPE)
quat1 /= np.linalg.norm(quat1)
pos2 = np.array((2.0, 3.0, 4.0))
quat2 = np.array((2.0, 3.0, 4.0, 5.0))
pos2 = np.array((2.0, 3.0, 4.0), dtype=DTYPE)
quat2 = np.array((2.0, 3.0, 4.0, 5.0), dtype=DTYPE)
quat2 /= np.linalg.norm(quat2)
control = np.hstack((pos1, pos2, quat1, quat2))
@@ -702,8 +718,8 @@ class MuJoCoRolloutTest(parameterized.TestCase):
nbatch = 1
nstep = 3
initial_state = np.zeros((nbatch, nstate))
ctrl = np.zeros((nbatch, nstep, model.nu))
initial_state = zeros(nbatch, nstate)
ctrl = zeros(nbatch, nstep, model.nu)
model.opt.solver = 10 # invalid solver type
with self.assertRaisesWithLiteralMatch(
@@ -718,7 +734,7 @@ class MuJoCoRolloutTest(parameterized.TestCase):
nbatch = 1
initial_state = np.zeros((nbatch, nstate))
initial_state = zeros(nbatch, nstate)
control = 'string'
with self.assertRaisesWithLiteralMatch(
@@ -726,7 +742,7 @@ class MuJoCoRolloutTest(parameterized.TestCase):
):
rollout.rollout(model, data, initial_state, control)
control = np.zeros((2, 3, 4, 5))
control = zeros(2, 3, 4, 5)
with self.assertRaisesWithLiteralMatch(
ValueError, 'control can have at most 3 dimensions'
):
@@ -740,28 +756,28 @@ class MuJoCoRolloutTest(parameterized.TestCase):
nbatch = 1
nstep = 3
initial_state = np.random.randn(nbatch, nstate + 1)
initial_state = randn(nbatch, nstate + 1)
with self.assertRaisesWithLiteralMatch(
ValueError, 'trailing dimension of initial_state must be 6, got 7'
):
rollout.rollout(model, data, initial_state)
initial_state = np.random.randn(nbatch, nstate)
control = np.random.randn(1, nstep, model.nu + 1)
initial_state = randn(nbatch, nstate)
control = randn(1, nstep, model.nu + 1)
with self.assertRaisesWithLiteralMatch(
ValueError, 'trailing dimension of control must be 2, got 3'
):
rollout.rollout(model, data, initial_state, control)
control = np.random.randn(nbatch, nstep, model.nu)
state = np.random.randn(nbatch, nstep + 1, nstate) # incompatible nstep
control = randn(nbatch, nstep, model.nu)
state = randn(nbatch, nstep + 1, nstate) # incompatible nstep
with self.assertRaisesWithLiteralMatch(
ValueError, 'dimension 1 inferred as 3 but state has 4'
):
rollout.rollout(model, data, initial_state, control, state=state)
initial_state = np.random.randn(nbatch, nstate)
control = np.random.randn(nbatch, nstep, model.nu)
initial_state = randn(nbatch, nstate)
control = randn(nbatch, nstep, model.nu)
bad_spec = mujoco.mjtState.mjSTATE_ACT
with self.assertRaisesWithLiteralMatch(
ValueError, 'control_spec can only contain bits in mjSTATE_USER'
@@ -776,8 +792,8 @@ class MuJoCoRolloutTest(parameterized.TestCase):
data = mujoco.MjData(model)
# step with a clean mjData
initial_state = np.random.randn(nstate)
control = np.random.randn(3, 3, model.nu)
initial_state = randn(nstate)
control = randn(3, 3, model.nu)
state, sensordata = rollout.rollout(model, data, initial_state, control)
# fill user fields with random values
@@ -788,7 +804,7 @@ class MuJoCoRolloutTest(parameterized.TestCase):
'mocap_pos',
'mocap_quat',
]:
setattr(data, attr, np.random.randn(*getattr(data, attr).shape))
setattr(data, attr, randn(*getattr(data, attr).shape))
# roll out again
state2, sensordata2 = rollout.rollout(model, data, initial_state, control)
@@ -802,8 +818,8 @@ class MuJoCoRolloutTest(parameterized.TestCase):
nstate = mujoco.mj_stateSize(model, mujoco.mjtState.mjSTATE_FULLPHYSICS)
data = mujoco.MjData(model)
initial_state = np.random.randn(nstate)
control = np.random.randn(3, 3, model.nu)
initial_state = randn(nstate)
control = randn(3, 3, model.nu)
state, sensordata = rollout.rollout(model, data, initial_state, control)
state2, sensordata2 = rollout.rollout([model], data, initial_state, control)
@@ -817,8 +833,8 @@ class MuJoCoRolloutTest(parameterized.TestCase):
nstate = mujoco.mj_stateSize(model, mujoco.mjtState.mjSTATE_FULLPHYSICS)
data = mujoco.MjData(model)
initial_state = np.random.randn(nstate)
control = np.random.randn(3, 3, model.nu)
initial_state = randn(nstate)
control = randn(3, 3, model.nu)
# Test passing empty lists for data
with self.assertRaisesWithLiteralMatch(
@@ -852,7 +868,7 @@ class MuJoCoRolloutTest(parameterized.TestCase):
):
with rollout.Rollout(nthread=0) as rollout_:
rollout_.rollout(
model, [copy.copy(data) for i in range(2)], initial_state, control
model, [copy.copy(data) for _ in range(2)], initial_state, control
)
with self.assertRaisesWithLiteralMatch(
@@ -872,7 +888,7 @@ class MuJoCoRolloutTest(parameterized.TestCase):
):
with rollout.Rollout(nthread=2) as rollout_:
rollout_.rollout(
model, [copy.copy(data) for i in range(3)], initial_state, control
model, [copy.copy(data) for _ in range(3)], initial_state, control
)
@absltest.skip(reason='Takes a long time to run')
@@ -887,7 +903,7 @@ class MuJoCoRolloutTest(parameterized.TestCase):
nstep = ((2**31) // (nstate * nbatch)) + 2
assert nstep * nstate * nbatch > 2**31
initial_state = np.random.randn(nbatch, nstate)
initial_state = randn(nbatch, nstate)
rollout.rollout(
model,
[copy.copy(data) for _ in range(nthread)],
@@ -901,7 +917,7 @@ class MuJoCoRolloutTest(parameterized.TestCase):
def get_state(model, data):
nstate = mujoco.mj_stateSize(model, mujoco.mjtState.mjSTATE_FULLPHYSICS)
state = np.empty(nstate)
state = empty(nstate)
mujoco.mj_getState(model, data, state, mujoco.mjtState.mjSTATE_FULLPHYSICS)
return state.reshape((1, nstate))
@@ -925,8 +941,8 @@ def one_rollout(
):
nstep = control.shape[0]
nstate = mujoco.mj_stateSize(model, mujoco.mjtState.mjSTATE_FULLPHYSICS)
state = np.empty((nstep, nstate))
sensordata = np.empty((nstep, model.nsensordata))
state = empty(nstep, nstate)
sensordata = empty(nstep, model.nsensordata)
mujoco.mj_resetData(model, data)
for t in range(nstep):
@@ -940,7 +956,7 @@ def ensure_2d(arg):
if arg is None:
return None
else:
return np.ascontiguousarray(np.atleast_2d(arg), dtype=np.float64)
return np.ascontiguousarray(np.atleast_2d(arg), dtype=DTYPE)
def ensure_3d(arg):
@@ -954,7 +970,7 @@ def ensure_3d(arg):
arg = arg[np.newaxis, np.newaxis, ...]
elif arg.ndim == 2:
arg = arg[np.newaxis, ...]
return np.ascontiguousarray(arg, dtype=np.float64)
return np.ascontiguousarray(arg, dtype=DTYPE)
def py_rollout(
@@ -974,8 +990,8 @@ def py_rollout(
nstate = mujoco.mj_stateSize(model[0], mujoco.mjtState.mjSTATE_FULLPHYSICS)
state = np.empty((nbatch, nstep, nstate))
sensordata = np.empty((nbatch, nstep, model[0].nsensordata))
state = empty(nbatch, nstep, nstate)
sensordata = empty(nbatch, nstep, model[0].nsensordata)
for r in range(nbatch):
state_r, sensordata_r = one_rollout(
model[r], data, initial_state[r], control[r], control_spec