Rename rollout variable nroll to nbatch.

PiperOrigin-RevId: 718196100
Change-Id: Ic36b1a96ea1af2de351539115d4eee051d195aaf
This commit is contained in:
Yuval Tassa
2025-01-21 21:03:08 -08:00
committed by Copybara-Service
parent ac76540eb8
commit 0090e1e908
5 changed files with 163 additions and 160 deletions
+89 -89
View File
@@ -185,11 +185,11 @@ class MuJoCoRolloutTest(parameterized.TestCase):
nstate = mujoco.mj_stateSize(model, mujoco.mjtState.mjSTATE_FULLPHYSICS)
data = mujoco.MjData(model)
nroll = 5 # number of rollouts
nbatch = 5 # number of rollouts
nstep = 1 # number of steps
initial_state = np.random.randn(nroll, nstate)
control = np.random.randn(nroll, nstep, model.nu)
initial_state = np.random.randn(nbatch, nstate)
control = np.random.randn(nbatch, nstep, model.nu)
state, sensordata = rollout.rollout(model, data, initial_state, control)
mujoco.mj_resetData(model, data)
@@ -198,108 +198,108 @@ class MuJoCoRolloutTest(parameterized.TestCase):
np.testing.assert_array_equal(sensordata, py_sensordata)
@parameterized.parameters(ALL_MODELS.keys())
def test_infer_nroll_initial_state(self, model_name):
def test_infer_nbatch_initial_state(self, model_name):
model = mujoco.MjModel.from_xml_string(ALL_MODELS[model_name])
nstate = mujoco.mj_stateSize(model, mujoco.mjtState.mjSTATE_FULLPHYSICS)
data = mujoco.MjData(model)
nroll = 5 # number of rollouts
nbatch = 5 # number of rollouts
nstep = 1 # number of steps
initial_state = np.random.randn(nroll, nstate)
initial_state = np.random.randn(nbatch, nstate)
control = np.random.randn(nstep, model.nu)
state, sensordata = rollout.rollout(model, data, initial_state, control)
mujoco.mj_resetData(model, data)
control = np.tile(control, (nroll, 1, 1))
control = np.tile(control, (nbatch, 1, 1))
py_state, py_sensordata = py_rollout(model, data, initial_state, control)
np.testing.assert_array_equal(state, py_state)
np.testing.assert_array_equal(sensordata, py_sensordata)
@parameterized.parameters(ALL_MODELS.keys())
def test_infer_nroll_control(self, model_name):
def test_infer_nbatch_control(self, model_name):
model = mujoco.MjModel.from_xml_string(ALL_MODELS[model_name])
nstate = mujoco.mj_stateSize(model, mujoco.mjtState.mjSTATE_FULLPHYSICS)
data = mujoco.MjData(model)
nroll = 5 # number of rollouts
nbatch = 5 # number of rollouts
nstep = 1 # number of steps
initial_state = np.random.randn(nstate)
control = np.random.randn(nroll, nstep, model.nu)
control = np.random.randn(nbatch, nstep, model.nu)
state, sensordata = rollout.rollout(model, data, initial_state, control)
mujoco.mj_resetData(model, data)
initial_state = np.tile(initial_state, (nroll, 1))
initial_state = np.tile(initial_state, (nbatch, 1))
py_state, py_sensordata = py_rollout(model, data, initial_state, control)
np.testing.assert_array_equal(state, py_state)
np.testing.assert_array_equal(sensordata, py_sensordata)
@parameterized.parameters(ALL_MODELS.keys())
def test_infer_nroll_warmstart(self, model_name):
def test_infer_nbatch_warmstart(self, model_name):
model = mujoco.MjModel.from_xml_string(ALL_MODELS[model_name])
nstate = mujoco.mj_stateSize(model, mujoco.mjtState.mjSTATE_FULLPHYSICS)
data = mujoco.MjData(model)
nroll = 5 # number of rollouts
nbatch = 5 # number of rollouts
nstep = 1 # number of steps
initial_state = np.random.randn(nstate)
control = np.random.randn(nstep, model.nu)
initial_warmstart = np.tile(data.qacc_warmstart.copy(), (nroll, 1))
initial_warmstart = np.tile(data.qacc_warmstart.copy(), (nbatch, 1))
state, sensordata = rollout.rollout(
model, data, initial_state, control, initial_warmstart=initial_warmstart
)
mujoco.mj_resetData(model, data)
initial_state = np.tile(initial_state, (nroll, 1))
control = np.tile(control, (nroll, 1, 1))
initial_state = np.tile(initial_state, (nbatch, 1))
control = np.tile(control, (nbatch, 1, 1))
py_state, py_sensordata = py_rollout(model, data, initial_state, control)
np.testing.assert_array_equal(state, py_state)
np.testing.assert_array_equal(sensordata, py_sensordata)
@parameterized.parameters(ALL_MODELS.keys())
def test_infer_nroll_state(self, model_name):
def test_infer_nbatch_state(self, model_name):
model = mujoco.MjModel.from_xml_string(ALL_MODELS[model_name])
nstate = mujoco.mj_stateSize(model, mujoco.mjtState.mjSTATE_FULLPHYSICS)
data = mujoco.MjData(model)
nroll = 5 # number of rollouts
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((nroll, nstep, nstate))
state = np.empty((nbatch, nstep, nstate))
state, sensordata = rollout.rollout(
model, data, initial_state, control, state=state
)
mujoco.mj_resetData(model, data)
initial_state = np.tile(initial_state, (nroll, 1))
control = np.tile(control, (nroll, 1, 1))
initial_state = np.tile(initial_state, (nbatch, 1))
control = np.tile(control, (nbatch, 1, 1))
py_state, py_sensordata = py_rollout(model, data, initial_state, control)
np.testing.assert_array_equal(state, py_state)
np.testing.assert_array_equal(sensordata, py_sensordata)
@parameterized.parameters(ALL_MODELS.keys())
def test_infer_nroll_sensordata(self, model_name):
def test_infer_nbatch_sensordata(self, model_name):
model = mujoco.MjModel.from_xml_string(ALL_MODELS[model_name])
nstate = mujoco.mj_stateSize(model, mujoco.mjtState.mjSTATE_FULLPHYSICS)
data = mujoco.MjData(model)
nroll = 5 # number of rollouts
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((nroll, nstep, model.nsensordata))
sensordata = np.empty((nbatch, nstep, model.nsensordata))
state, sensordata = rollout.rollout(
model, data, initial_state, control, sensordata=sensordata
)
mujoco.mj_resetData(model, data)
initial_state = np.tile(initial_state, (nroll, 1))
control = np.tile(control, (nroll, 1, 1))
initial_state = np.tile(initial_state, (nbatch, 1))
control = np.tile(control, (nbatch, 1, 1))
py_state, py_sensordata = py_rollout(model, data, initial_state, control)
np.testing.assert_array_equal(state, py_state)
np.testing.assert_array_equal(sensordata, py_sensordata)
@@ -310,13 +310,13 @@ class MuJoCoRolloutTest(parameterized.TestCase):
nstate = mujoco.mj_stateSize(model, mujoco.mjtState.mjSTATE_FULLPHYSICS)
data = mujoco.MjData(model)
nroll = 1 # number of rollouts
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((nroll, nstep, nstate))
sensordata = np.empty((nroll, nstep, model.nsensordata))
state = np.empty((nbatch, nstep, nstate))
sensordata = np.empty((nbatch, nstep, model.nsensordata))
rollout.rollout(
model, data, initial_state, control, state=state, sensordata=sensordata
)
@@ -332,11 +332,11 @@ class MuJoCoRolloutTest(parameterized.TestCase):
nstate = mujoco.mj_stateSize(model, mujoco.mjtState.mjSTATE_FULLPHYSICS)
data = mujoco.MjData(model)
nroll = 2 # number of initial states
nbatch = 2 # number of initial states
nstep = 3 # number of timesteps
initial_state = np.random.randn(nroll, nstate)
control = np.random.randn(nroll, nstep, model.nu)
initial_state = np.random.randn(nbatch, nstate)
control = np.random.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)
@@ -345,26 +345,26 @@ class MuJoCoRolloutTest(parameterized.TestCase):
@parameterized.parameters(ALL_MODELS.keys())
def test_multi_model(self, model_name):
nroll = 3 # number of initial states and models
nbatch = 3 # number of initial states and models
nstep = 3 # number of timesteps
spec = mujoco.MjSpec.from_string(ALL_MODELS[model_name])
if len(spec.bodies) > 1:
model = []
for i in range(nroll):
for i in range(nbatch):
body = spec.bodies[1]
assert body.name != 'world'
body.pos = body.pos + i
model.append(spec.compile())
else:
model = [spec.compile() for _ in range(nroll)]
model = [spec.compile() for _ in range(nbatch)]
nstate = mujoco.mj_stateSize(model[0], mujoco.mjtState.mjSTATE_FULLPHYSICS)
data = mujoco.MjData(model[0])
initial_state = np.random.randn(nroll, nstate)
control = np.random.randn(nroll, nstep, model[0].nu)
initial_state = np.random.randn(nbatch, nstate)
control = np.random.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)
@@ -377,12 +377,12 @@ class MuJoCoRolloutTest(parameterized.TestCase):
nstate = mujoco.mj_stateSize(model, mujoco.mjtState.mjSTATE_FULLPHYSICS)
data = mujoco.MjData(model)
nroll = 2 # number of rollouts
nbatch = 2 # number of rollouts
nstep = 3 # number of timesteps
initial_state = np.random.randn(nroll, nstate)
control = np.random.randn(nroll, 1, model.nu)
state = np.empty((nroll, nstep, nstate))
initial_state = np.random.randn(nbatch, nstate)
control = np.random.randn(nbatch, 1, model.nu)
state = np.empty((nbatch, nstep, nstate))
state, sensordata = rollout.rollout(
model, data, initial_state, control, state=state
)
@@ -398,10 +398,10 @@ class MuJoCoRolloutTest(parameterized.TestCase):
nstate = mujoco.mj_stateSize(model, mujoco.mjtState.mjSTATE_FULLPHYSICS)
data = mujoco.MjData(model)
nroll = 4 # number of rollouts
nbatch = 4 # number of rollouts
nstep = 3 # number of timesteps
initial_state = np.random.randn(nroll, nstate)
initial_state = np.random.randn(nbatch, nstate)
control_spec = (
mujoco.mjtState.mjSTATE_CTRL
@@ -409,7 +409,7 @@ class MuJoCoRolloutTest(parameterized.TestCase):
| mujoco.mjtState.mjSTATE_XFRC_APPLIED
)
ncontrol = mujoco.mj_stateSize(model, control_spec)
control = np.random.randn(nroll, nstep, ncontrol)
control = np.random.randn(nbatch, nstep, ncontrol)
state, sensordata = rollout.rollout(
model, data, initial_state, control, control_spec=control_spec
@@ -426,8 +426,8 @@ class MuJoCoRolloutTest(parameterized.TestCase):
nstate = mujoco.mj_stateSize(model, mujoco.mjtState.mjSTATE_FULLPHYSICS)
data = mujoco.MjData(model)
nroll = 4 # number of rollouts
initial_state = np.empty((nroll, nstate))
nbatch = 4 # number of rollouts
initial_state = np.empty((nbatch, nstate))
# get diverging (0, 2) and non-diverging (1, 3) states
mujoco.mj_getState(
@@ -446,7 +446,7 @@ class MuJoCoRolloutTest(parameterized.TestCase):
nstep = 10000 # divergence after ~15s, timestep = 2e-3
state = np.random.randn(nroll, nstep, nstate)
state = np.random.randn(nbatch, nstep, nstate)
rollout.rollout(model, data, initial_state, state=state)
@@ -464,19 +464,19 @@ class MuJoCoRolloutTest(parameterized.TestCase):
model = mujoco.MjModel.from_xml_string(TEST_XML)
nstate = mujoco.mj_stateSize(model, mujoco.mjtState.mjSTATE_FULLPHYSICS)
num_workers = 32
nroll = 100
nbatch = 100
nstep = 5
initial_state = np.random.randn(nroll, nstate)
state = np.empty((nroll, nstep, nstate))
sensordata = np.empty((nroll, nstep, model.nsensordata))
control = np.random.randn(nroll, nstep, model.nu)
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)
thread_local = threading.local()
def thread_initializer():
thread_local.data = mujoco.MjData(model)
model_list = [copy.copy(model) for _ in range(nroll)]
model_list = [copy.copy(model) for _ in range(nbatch)]
def call_rollout(initial_state, control, state, sensordata):
rollout.rollout(
@@ -490,7 +490,7 @@ class MuJoCoRolloutTest(parameterized.TestCase):
sensordata=sensordata,
)
n = nroll // num_workers # integer division
n = nbatch // num_workers # integer division
chunks = [] # a list of tuples, one per worker
for i in range(num_workers - 1):
chunks.append((
@@ -526,14 +526,14 @@ class MuJoCoRolloutTest(parameterized.TestCase):
model = mujoco.MjModel.from_xml_string(TEST_XML)
nstate = mujoco.mj_stateSize(model, mujoco.mjtState.mjSTATE_FULLPHYSICS)
num_workers = 32
nroll = 100
nbatch = 100
nstep = 5
initial_state = np.random.randn(nroll, nstate)
state = np.empty((nroll, nstep, nstate))
sensordata = np.empty((nroll, nstep, model.nsensordata))
control = np.random.randn(nroll, nstep, model.nu)
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)
model_list = [copy.copy(model) for _ in range(nroll)]
model_list = [copy.copy(model) for _ in range(nbatch)]
data_list = [mujoco.MjData(model) for _ in range(num_workers)]
rollout.rollout(
@@ -555,14 +555,14 @@ class MuJoCoRolloutTest(parameterized.TestCase):
model = mujoco.MjModel.from_xml_string(TEST_XML)
nstate = mujoco.mj_stateSize(model, mujoco.mjtState.mjSTATE_FULLPHYSICS)
num_workers = 32
nroll = 100
nbatch = 100
nstep = 5
initial_state = np.random.randn(nroll, nstate)
state = np.empty((nroll, nstep, nstate))
sensordata = np.empty((nroll, nstep, model.nsensordata))
control = np.random.randn(nroll, nstep, model.nu)
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)
model_list = [copy.copy(model) for _ in range(nroll)]
model_list = [copy.copy(model) for _ in range(nbatch)]
data_list = [mujoco.MjData(model) for _ in range(num_workers)]
with rollout.Rollout(nthread=num_workers) as rollout_:
@@ -604,14 +604,14 @@ class MuJoCoRolloutTest(parameterized.TestCase):
model = mujoco.MjModel.from_xml_string(TEST_XML)
nstate = mujoco.mj_stateSize(model, mujoco.mjtState.mjSTATE_FULLPHYSICS)
num_workers = 32
nroll = 100
nbatch = 100
nstep = 5
initial_state = np.random.randn(nroll, nstate)
state = np.empty((nroll, nstep, nstate))
sensordata = np.empty((nroll, nstep, model.nsensordata))
control = np.random.randn(nroll, nstep, model.nu)
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)
model_list = [copy.copy(model) for _ in range(nroll)]
model_list = [copy.copy(model) for _ in range(nbatch)]
data_list = [mujoco.MjData(model) for _ in range(num_workers)]
for _ in range(2):
@@ -699,11 +699,11 @@ class MuJoCoRolloutTest(parameterized.TestCase):
nstate = mujoco.mj_stateSize(model, mujoco.mjtState.mjSTATE_FULLPHYSICS)
data = mujoco.MjData(model)
nroll = 1
nbatch = 1
nstep = 3
initial_state = np.zeros((nroll, nstate))
ctrl = np.zeros((nroll, nstep, model.nu))
initial_state = np.zeros((nbatch, nstate))
ctrl = np.zeros((nbatch, nstep, model.nu))
model.opt.solver = 10 # invalid solver type
with self.assertRaisesWithLiteralMatch(
@@ -716,9 +716,9 @@ class MuJoCoRolloutTest(parameterized.TestCase):
nstate = mujoco.mj_stateSize(model, mujoco.mjtState.mjSTATE_FULLPHYSICS)
data = mujoco.MjData(model)
nroll = 1
nbatch = 1
initial_state = np.zeros((nroll, nstate))
initial_state = np.zeros((nbatch, nstate))
control = 'string'
with self.assertRaisesWithLiteralMatch(
@@ -737,31 +737,31 @@ class MuJoCoRolloutTest(parameterized.TestCase):
nstate = mujoco.mj_stateSize(model, mujoco.mjtState.mjSTATE_FULLPHYSICS)
data = mujoco.MjData(model)
nroll = 1
nbatch = 1
nstep = 3
initial_state = np.random.randn(nroll, nstate + 1)
initial_state = np.random.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(nroll, nstate)
initial_state = np.random.randn(nbatch, nstate)
control = np.random.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(nroll, nstep, model.nu)
state = np.random.randn(nroll, nstep + 1, nstate) # incompatible nstep
control = np.random.randn(nbatch, nstep, model.nu)
state = np.random.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(nroll, nstate)
control = np.random.randn(nroll, nstep, model.nu)
initial_state = np.random.randn(nbatch, nstate)
control = np.random.randn(nbatch, nstep, model.nu)
bad_spec = mujoco.mjtState.mjSTATE_ACT
with self.assertRaisesWithLiteralMatch(
ValueError, 'control_spec can only contain bits in mjSTATE_USER'
@@ -946,17 +946,17 @@ def py_rollout(
):
initial_state = ensure_2d(initial_state)
control = ensure_3d(control)
nroll = initial_state.shape[0]
nbatch = initial_state.shape[0]
nstep = control.shape[1]
if isinstance(model, mujoco.MjModel):
model = [copy.copy(model) for _ in range(nroll)]
model = [copy.copy(model) for _ in range(nbatch)]
nstate = mujoco.mj_stateSize(model[0], mujoco.mjtState.mjSTATE_FULLPHYSICS)
state = np.empty((nroll, nstep, nstate))
sensordata = np.empty((nroll, nstep, model[0].nsensordata))
for r in range(nroll):
state = np.empty((nbatch, nstep, nstate))
sensordata = np.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
)