Copybara import of the project:

--
3a95b62f59e81bfef0f076afb173ecc14b27943d by Levi Burner <leviburner@gmail.com>:

rollout prototype native threadpool for comparing to python threads

--
efd8be1124ac839b902de45973a3ca8b9f2215e6 by Levi Burner <leviburner@gmail.com>:

copy mjpcs threadpool into python bindings

--
75603eea3e8362e354a9675e8a6cd14e56ec3d28 by Levi Burner <leviburner@gmail.com>:

rollout use threadpool as translation unit

--
06b90febd021663f6cc81fd7895e4d6e2008ed97 by Levi Burner <leviburner@gmail.com>:

rollout add chunk_divisor parameter

--
298ab2f3c0d6e12530832c3cdbf784dd92d54806 by Levi Burner <leviburner@gmail.com>:

rollout add native threading test

--
169cf9978e7abad6edd1392b8e6aab995e4f8f10 by Levi Burner <leviburner@gmail.com>:

rollout exchange chunk_divisor arg for chunk_size

--
265af851d74432d261277d3dbda11cdef1841bc8 by Levi Burner <leviburner@gmail.com>:

rollout fix cosmetics

--
1e8bffa88bf36190501b334bef31147e23db39f7 by Levi Burner <leviburner@gmail.com>:

make native rollout a class instead of a function

--
ba788214b047577f58c41ce0ab6c62c277cd8b0d by Levi Burner <leviburner@gmail.com>:

rollout update docs and changelog

--
e4cb7732319e04cba2ab2c2ad848c659f6309808 by Levi Burner <leviburner@gmail.com>:

rollout don't register atexit handler for Rollout objects

--
5a08d2efdbbbb01d4b1231ff9a36a1dc44f4d9ee by Levi Burner <leviburner@gmail.com>:

rollout nthread kwarg, rename shutdown_pool to close, fixups

--
f622378543596a208339af0208fa3a70bf2a8007 by Levi Burner <leviburner@gmail.com>:

rollout add missing .close() calls

--
50f3ebca43c53eac03f03943c34bb1e46967bd4f by Levi Burner <leviburner@gmail.com>:

rollout return immediately

COPYBARA_INTEGRATE_REVIEW=https://github.com/google-deepmind/mujoco/pull/2282 from aftersomemath:rollout-threaded 50f3ebca43c53eac03f03943c34bb1e46967bd4f
PiperOrigin-RevId: 706744277
Change-Id: I1ab2263b7d6ce30cf1908aec8fd5f2eb976a19e6
This commit is contained in:
Levi Burner
2024-12-16 09:56:12 -08:00
committed by Copybara-Service
parent b26d6f0466
commit a7eb6efd4e
8 changed files with 761 additions and 203 deletions
+113 -3
View File
@@ -355,7 +355,7 @@ class MuJoCoRolloutTest(parameterized.TestCase):
body.pos = body.pos + i
model.append(spec.compile())
else:
model = [spec.compile() for i in range(nroll)]
model = [spec.compile() for _ in range(nroll)]
nstate = mujoco.mj_stateSize(model[0], mujoco.mjtState.mjSTATE_FULLPHYSICS)
data = mujoco.MjData(model[0])
@@ -461,7 +461,7 @@ 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 = 10000
nroll = 100
nstep = 5
initial_state = np.random.randn(nroll, nstate)
state = np.empty((nroll, nstep, nstate))
@@ -478,7 +478,7 @@ class MuJoCoRolloutTest(parameterized.TestCase):
def call_rollout(initial_state, control, state, sensordata):
rollout.rollout(
model_list,
thread_local.data,
[thread_local.data],
initial_state,
control,
skip_checks=True,
@@ -519,6 +519,116 @@ class MuJoCoRolloutTest(parameterized.TestCase):
np.testing.assert_array_equal(state, py_state)
np.testing.assert_array_equal(sensordata, py_sensordata)
def test_threading_native(self):
model = mujoco.MjModel.from_xml_string(TEST_XML)
nstate = mujoco.mj_stateSize(model, mujoco.mjtState.mjSTATE_FULLPHYSICS)
num_workers = 32
nroll = 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)
model_list = [model] * nroll
data_list = [mujoco.MjData(model) for _ in range(num_workers)]
rollout.rollout(
model_list,
data_list,
initial_state,
control,
nstep=nstep,
state=state,
sensordata=sensordata,
)
data = mujoco.MjData(model)
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)
def test_threading_native_persistent_object(self):
model = mujoco.MjModel.from_xml_string(TEST_XML)
nstate = mujoco.mj_stateSize(model, mujoco.mjtState.mjSTATE_FULLPHYSICS)
num_workers = 32
nroll = 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)
model_list = [model] * nroll
data_list = [mujoco.MjData(model) for _ in range(num_workers)]
with rollout.Rollout(nthread=num_workers) as rollout_:
for _ in range(2):
rollout_.rollout(
model_list,
data_list,
initial_state,
control,
nstep=nstep,
state=state,
sensordata=sensordata,
)
data = mujoco.MjData(model)
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)
rollout_ = rollout.Rollout(nthread=num_workers)
for _ in range(2):
rollout_.rollout(
model_list,
data_list,
initial_state,
control,
nstep=nstep,
state=state,
sensordata=sensordata,
)
data = mujoco.MjData(model)
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)
rollout_.close()
def test_threading_native_persistent_function(self):
model = mujoco.MjModel.from_xml_string(TEST_XML)
nstate = mujoco.mj_stateSize(model, mujoco.mjtState.mjSTATE_FULLPHYSICS)
num_workers = 32
nroll = 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)
model_list = [model] * nroll
data_list = [mujoco.MjData(model) for _ in range(num_workers)]
for _ in range(2):
rollout.rollout(
model_list,
data_list,
initial_state,
control,
nstep=nstep,
state=state,
sensordata=sensordata,
persistent_pool=True,
)
data = mujoco.MjData(model)
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)
rollout.shutdown_persistent_pool()
# ---------------------------- test advanced operation
def test_warmstart(self):