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:
committed by
Copybara-Service
parent
116cfe9c70
commit
a0f014593b
@@ -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;
|
||||
|
||||
@@ -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
@@ -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
|
||||
|
||||
Reference in New Issue
Block a user