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
+1 -1
View File
@@ -383,7 +383,7 @@ target_link_libraries(
structs_header
)
mujoco_pybind11_module(_rollout rollout.cc)
mujoco_pybind11_module(_rollout rollout.cc threadpool.cc)
target_link_libraries(_rollout PRIVATE functions_header mujoco raw)
mujoco_pybind11_module(
+162 -64
View File
@@ -13,6 +13,7 @@
// limitations under the License.
#include <iostream>
#include <memory>
#include <optional>
#include <sstream>
@@ -20,6 +21,7 @@
#include "errors.h"
#include "raw.h"
#include "structs.h"
#include "threadpool.h"
#include <pybind11/buffer_info.h>
#include <pybind11/numpy.h>
#include <pybind11/pybind11.h>
@@ -31,14 +33,24 @@ namespace {
namespace py = ::pybind11;
using PyCArray = py::array_t<mjtNum, py::array::c_style>;
// NOLINTBEGIN(whitespace/line_length)
const auto rollout_init_doc = R"(
Construct a rollout object containing a thread pool for parallel rollouts.
input arguments (optional):
nthread integer, number of threads in pool
if zero, this pool is not started and rollouts run on the calling thread
)";
const auto rollout_doc = R"(
Roll out open-loop trajectories from initial states, get resulting states and sensor values.
input arguments (required):
model list of MjModel instances of length nroll
data associated instance of MjData
data list of associated MjData instances of length nthread
nstep integer, number of steps to be taken for each trajectory
control_spec specification of controls, ncontrol = mj_stateSize(m, control_spec)
state0 (nroll x nstate) nroll initial state vectors,
@@ -49,12 +61,14 @@ Roll out open-loop trajectories from initial states, get resulting states and se
output arguments (optional):
state (nroll x nstep x nstate) nroll nstep states
sensordata (nroll x nstep x nsendordata) nroll trajectories of nstep sensordata vectors
chunk_size integer, determines threadpool chunk size. If unspecified
chunk_size = max(1, nroll / (nthread * 10))
)";
// C-style rollout function, assumes all arguments are valid
// all input fields of d are initialised, contents at call time do not matter
// after returning, d will contain the last step of the last rollout
void _unsafe_rollout(std::vector<const mjModel*>& m, mjData* d, int nroll, int nstep, unsigned int control_spec,
void _unsafe_rollout(std::vector<const mjModel*>& m, mjData* d, int start_roll, int end_roll, int nstep, unsigned int control_spec,
const mjtNum* state0, const mjtNum* warmstart0, const mjtNum* control,
mjtNum* state, mjtNum* sensordata) {
// sizes
@@ -75,7 +89,7 @@ void _unsafe_rollout(std::vector<const mjModel*>& m, mjData* d, int nroll, int n
}
// loop over rollouts
for (int r = 0; r < nroll; r++) {
for (int r = start_roll; r < end_roll; r++) {
// clear user inputs if unspecified
if (!(control_spec & mjSTATE_MOCAP_POS)) {
for (int i = 0; i < nbody; i++) {
@@ -158,6 +172,43 @@ void _unsafe_rollout(std::vector<const mjModel*>& m, mjData* d, int nroll, int n
}
}
// C-style threaded version of _unsafe_rollout
void _unsafe_rollout_threaded(std::vector<const mjModel*>& m, std::vector<mjData*>& d,
int nroll, 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 = nroll / chunk_size;
int chunk_remainder = nroll % chunk_size;
int njobs = (chunk_remainder > 0) ? nfulljobs + 1 : nfulljobs;
// Reset the pool counter
pool->ResetCount();
// schedule all jobs of full (chunk) size
for (int j = 0; j < nfulljobs; j++) {
auto task = [=, &m, &d](void) {
int id = pool->WorkerId();
_unsafe_rollout(m, d[id], j*chunk_size, (j+1)*chunk_size,
nstep, control_spec, state0, warmstart0, control, state, sensordata);
};
pool->Schedule(task);
}
// schedule any remaining jobs of size < chunk_size
if (chunk_remainder > 0) {
auto task = [=, &m, &d](void) {
_unsafe_rollout(m, d[pool->WorkerId()], nfulljobs*chunk_size,
nfulljobs*chunk_size+chunk_remainder,
nstep, control_spec, state0, warmstart0, control, state, sensordata);
};
pool->Schedule(task);
}
// wait for job counter to incremented up to the number of jobs submitted by this thread
pool->WaitCount(njobs);
}
// NOLINTEND(whitespace/line_length)
// check size of optional argument to rollout(), return raw pointer
@@ -181,71 +232,118 @@ mjtNum* get_array_ptr(std::optional<const py::array_t<mjtNum>> arg,
return static_cast<mjtNum*>(info.ptr);
}
class Rollout {
public:
Rollout(int nthread) : nthread_(nthread) {
if (this->nthread_ > 0) {
this->pool_ = std::make_shared<ThreadPool>(this->nthread_);
}
}
void rollout(py::list m, py::list d, int nstep, unsigned int control_spec,
const PyCArray state0, std::optional<const PyCArray> warmstart0,
std::optional<const PyCArray> control,
std::optional<const PyCArray> state,
std::optional<const PyCArray> sensordata,
std::optional<int> chunk_size) {
// get raw pointers
int nroll = state0.shape(0);
std::vector<const raw::MjModel*> model_ptrs(nroll);
for (int r = 0; r < nroll; r++) {
model_ptrs[r] = m[r].cast<const MjModelWrapper*>()->get();
}
// check length d and nthread are consistent
if (this->nthread_ == 0 && py::len(d) > 1) {
std::ostringstream msg;
msg << "More than one data instance passed but "
<< "rollout is configured to run on main thread";
py::value_error(msg.str());
} else if (this->nthread_ != py::len(d)) {
std::ostringstream msg;
msg << "Length of data: " << py::len(d)
<< " not equal to nthread: " << this->nthread_;
py::value_error(msg.str());
}
std::vector<raw::MjData*> data_ptrs(py::len(d));
for (int t = 0; t < py::len(d); t++) {
data_ptrs[t] = d[t].cast<MjDataWrapper*>()->get();
}
// check that some steps need to be taken, return if not
if (nstep < 1) {
return;
}
// get sizes
int nstate = mj_stateSize(model_ptrs[0], mjSTATE_FULLPHYSICS);
int ncontrol = mj_stateSize(model_ptrs[0], control_spec);
mjtNum* state0_ptr = get_array_ptr(state0, "state0", nroll, 1, nstate);
mjtNum* warmstart0_ptr =
get_array_ptr(warmstart0, "warmstart0", nroll, 1, model_ptrs[0]->nv);
mjtNum* control_ptr =
get_array_ptr(control, "control", nroll, nstep, ncontrol);
mjtNum* state_ptr = get_array_ptr(state, "state", nroll, nstep, nstate);
mjtNum* sensordata_ptr = get_array_ptr(sensordata, "sensordata", nroll,
nstep, model_ptrs[0]->nsensordata);
// perform rollouts
{
// release the GIL
py::gil_scoped_release no_gil;
// call unsafe rollout function, multi or single threaded
if (this->nthread_ > 0 && nroll > 1) {
int chunk_size_final = 1;
if (!chunk_size.has_value()) {
chunk_size_final = std::max(1, nroll / (10 * this->nthread_));
} else {
chunk_size_final = *chunk_size;
}
InterceptMjErrors(_unsafe_rollout_threaded)(
model_ptrs, data_ptrs, nroll, nstep, control_spec, state0_ptr,
warmstart0_ptr, control_ptr, state_ptr, sensordata_ptr,
this->pool_.get(), chunk_size_final);
} else {
InterceptMjErrors(_unsafe_rollout)(
model_ptrs, data_ptrs[0], 0, nroll, nstep, control_spec, state0_ptr,
warmstart0_ptr, control_ptr, state_ptr, sensordata_ptr);
}
}
}
private:
int nthread_;
std::shared_ptr<ThreadPool> pool_;
};
PYBIND11_MODULE(_rollout, pymodule) {
namespace py = ::pybind11;
using PyCArray = py::array_t<mjtNum, py::array::c_style>;
// roll out open loop trajectories from multiple initial states
// get subsequent states and corresponding sensor values
pymodule.def(
"rollout",
[](py::list m, MjDataWrapper& d,
int nstep, unsigned int control_spec,
const PyCArray state0,
std::optional<const PyCArray> warmstart0,
std::optional<const PyCArray> control,
std::optional<const PyCArray> state,
std::optional<const PyCArray> sensordata
) {
// get raw pointers
int nroll = state0.shape(0);
std::vector<const raw::MjModel*> model_ptrs(nroll);
for (int r = 0; r < nroll; r++) {
model_ptrs[r] = m[r].cast<const MjModelWrapper*>()->get();
}
raw::MjData* data = d.get();
// check that some steps need to be taken, return if not
if (nstep < 1) {
return;
}
// get sizes
int nstate = mj_stateSize(model_ptrs[0], mjSTATE_FULLPHYSICS);
int ncontrol = mj_stateSize(model_ptrs[0], control_spec);
mjtNum* state0_ptr = get_array_ptr(state0, "state0", nroll, 1, nstate);
mjtNum* warmstart0_ptr = get_array_ptr(warmstart0, "warmstart0", nroll,
1, model_ptrs[0]->nv);
mjtNum* control_ptr = get_array_ptr(control, "control", nroll,
nstep, ncontrol);
mjtNum* state_ptr = get_array_ptr(state, "state", nroll, nstep, nstate);
mjtNum* sensordata_ptr = get_array_ptr(sensordata, "sensordata", nroll,
nstep, model_ptrs[0]->nsensordata);
// perform rollouts
{
// release the GIL
py::gil_scoped_release no_gil;
// call unsafe rollout function
InterceptMjErrors(_unsafe_rollout)(
model_ptrs, data, nroll, nstep, control_spec, state0_ptr,
warmstart0_ptr, control_ptr, state_ptr, sensordata_ptr);
}
},
py::arg("model"),
py::arg("data"),
py::arg("nstep"),
py::arg("control_spec"),
py::arg("state0"),
py::arg("warmstart0") = py::none(),
py::arg("control") = py::none(),
py::arg("state") = py::none(),
py::arg("sensordata") = py::none(),
py::doc(rollout_doc)
);
py::class_<Rollout>(pymodule, "Rollout")
.def(
py::init([](int nthread) {
return std::make_unique<Rollout>(nthread);
}),
py::kw_only(),
py::arg("nthread"),
py::doc(rollout_init_doc))
.def(
"rollout",
&Rollout::rollout,
py::arg("model"),
py::arg("data"),
py::arg("nstep"),
py::arg("control_spec"),
py::arg("state0"),
py::arg("warmstart0") = py::none(),
py::arg("control") = py::none(),
py::arg("state") = py::none(),
py::arg("sensordata") = py::none(),
py::arg("chunk_size") = py::none(),
py::doc(rollout_doc));
}
} // namespace
+273 -128
View File
@@ -14,6 +14,7 @@
# ==============================================================================
"""Roll out open-loop trajectories from initial states, get subsequent states and sensor values."""
import atexit
from collections.abc import Sequence
from typing import Optional, Union
@@ -23,9 +24,243 @@ import numpy as np
from numpy import typing as npt
class Rollout:
"""Rollout object containing a thread pool for parallel rollouts."""
def __init__(self, *, nthread: Optional[int] = None):
"""Construct a rollout object containing a thread pool for parallel rollouts.
Args:
nthread: Number of threads in pool.
If zero, this pool is not started and rollouts run on the calling thread.
""" # fmt: skip
self.nthread = 0 if nthread is None else nthread
self.rollout_ = _rollout.Rollout(nthread=self.nthread)
def __enter__(self):
return self
def __exit__(self, exc_type, exc_val, exc_tb):
self.close()
def close(self):
del self.rollout_
self.rollout_ = None
def rollout(
self,
model: Union[mujoco.MjModel, Sequence[mujoco.MjModel]],
data: Union[mujoco.MjData, Sequence[mujoco.MjData]],
initial_state: npt.ArrayLike,
control: Optional[npt.ArrayLike] = None,
*, # require subsequent arguments to be named
control_spec: int = mujoco.mjtState.mjSTATE_CTRL.value,
skip_checks: bool = False,
nstep: Optional[int] = None,
initial_warmstart: Optional[npt.ArrayLike] = None,
state: Optional[npt.ArrayLike] = None,
sensordata: Optional[npt.ArrayLike] = None,
chunk_size: Optional[int] = None,
):
"""Rolls out open-loop trajectories from initial states, get subsequent state and sensor values.
Python wrapper for rollout.cc, see documentation therein.
Infers nroll and nstep.
Tiles inputs with singleton dimensions.
Allocates outputs if none are given.
Args:
model: An instance or length nroll sequence of MjModel with the same size signature.
data: Associated mjData instance or sequence of instances with length nthread.
initial_state: Array of initial states from which to roll out trajectories.
([nroll or 1] x nstate)
control: Open-loop controls array to apply during the rollouts.
([nroll or 1] x [nstep or 1] x ncontrol)
control_spec: mjtState specification of control vectors.
skip_checks: Whether to skip internal shape and type checks.
nstep: Number of steps in rollouts (inferred if unspecified).
initial_warmstart: Initial qfrc_warmstart array (optional).
([nroll or 1] x nv)
state: State output array (optional).
(nroll x nstep x nstate)
sensordata: Sensor data output array (optional).
(nroll x nstep x nsensordata)
chunk_size: Determines threadpool chunk size. If unspecified,
chunk_size = max(1, nroll / (nthread * 10))
Returns:
state:
State output array, (nroll x nstep x nstate).
sensordata:
Sensor data output array, (nroll x nstep x nsensordata).
Raises:
RuntimeError: rollout requested after thread pool shutdown.
ValueError: bad shapes or sizes.
""" # fmt: skip
if self.rollout_ is None:
raise RuntimeError('rollout requested after thread pool shutdown')
# skip_checks shortcut:
# don't infer nroll/nstep
# don't support singleton expansion
# don't allocate output arrays
# just call rollout and return
if skip_checks:
self.rollout_.rollout(
model,
data,
nstep,
control_spec,
initial_state,
initial_warmstart,
control,
state,
sensordata,
chunk_size,
)
return state, sensordata
if not isinstance(model, mujoco.MjModel):
model = list(model)
# check control_spec
if control_spec & ~mujoco.mjtState.mjSTATE_USER.value:
raise ValueError('control_spec can only contain bits in mjSTATE_USER')
# check types
if nstep and not isinstance(nstep, int):
raise ValueError('nstep must be an integer')
if chunk_size and not isinstance(chunk_size, int):
raise ValueError('chunk_size must be an integer')
_check_must_be_numeric(
initial_state=initial_state,
initial_warmstart=initial_warmstart,
control=control,
state=state,
sensordata=sensordata,
)
# check number of dimensions
_check_number_of_dimensions(
2, initial_state=initial_state, initial_warmstart=initial_warmstart
)
_check_number_of_dimensions(
3, control=control, state=state, sensordata=sensordata
)
# ensure 2D, make contiguous, row-major (C ordering)
initial_state = _ensure_2d(initial_state)
initial_warmstart = _ensure_2d(initial_warmstart)
# ensure 3D, make contiguous, row-major (C ordering)
control = _ensure_3d(control)
state = _ensure_3d(state)
sensordata = _ensure_3d(sensordata)
# infer nroll, check for incompatibilities
nroll = _infer_dimension(
0,
1,
initial_state=initial_state,
initial_warmstart=initial_warmstart,
control=control,
state=state,
sensordata=sensordata,
)
if isinstance(model, list) and nroll == 1:
nroll = len(model)
if isinstance(model, list) and len(model) != nroll:
raise ValueError(
f'nroll inferred as {nroll} but model is length {len(model)}'
)
elif not isinstance(model, list):
model = [model] # Use a length 1 list to simplify code below
if not isinstance(data, list):
data = [data] # Use a length 1 list to simplify code below
# infer nstep, check for incompatibilities
nstep = _infer_dimension(
1, nstep or 1, control=control, state=state, sensordata=sensordata
)
# get nstate/ncontrol/nv/nsensordata
# check that they are equal across models
nstate = mujoco.mj_stateSize(
model[0], mujoco.mjtState.mjSTATE_FULLPHYSICS.value
)
ncontrol = mujoco.mj_stateSize(model[0], control_spec)
nv = model[0].nv
nsensordata = model[0].nsensordata
for m in model[1:]:
if (
nstate
!= mujoco.mj_stateSize(m, mujoco.mjtState.mjSTATE_FULLPHYSICS.value)
or ncontrol != mujoco.mj_stateSize(m, control_spec)
or nv != m.nv
or nsensordata != m.nsensordata
):
raise ValueError('models are not compatible')
# check trailing dimensions
_check_trailing_dimension(nstate, initial_state=initial_state, state=state)
_check_trailing_dimension(ncontrol, control=control)
_check_trailing_dimension(nv, initial_warmstart=initial_warmstart)
_check_trailing_dimension(nsensordata, sensordata=sensordata)
# tile input arrays/lists if required (singleton expansion)
model = model * nroll if len(model) == 1 else model
initial_state = _tile_if_required(initial_state, nroll)
initial_warmstart = _tile_if_required(initial_warmstart, nroll)
control = _tile_if_required(control, nroll, nstep)
# allocate output if not provided
if state is None:
state = np.empty((nroll, nstep, nstate))
if sensordata is None:
sensordata = np.empty((nroll, nstep, nsensordata))
# call rollout
self.rollout_.rollout(
model,
data,
nstep,
control_spec,
initial_state,
initial_warmstart,
control,
state,
sensordata,
chunk_size,
)
# return outputs
return state, sensordata
persistent_rollout = None
def shutdown_persistent_pool():
"""Shutdown the persistent thread pool that is optionally created by rollout.
This is called automatically interpreter shutdown, but can also be called manually.
""" # fmt: skip
global persistent_rollout
if persistent_rollout is not None:
persistent_rollout.close()
persistent_rollout = None
atexit.register(shutdown_persistent_pool)
def rollout(
model: Union[mujoco.MjModel, Sequence[mujoco.MjModel]],
data: mujoco.MjData,
data: Union[mujoco.MjData, Sequence[mujoco.MjData]],
initial_state: npt.ArrayLike,
control: Optional[npt.ArrayLike] = None,
*, # require subsequent arguments to be named
@@ -35,6 +270,8 @@ def rollout(
initial_warmstart: Optional[npt.ArrayLike] = None,
state: Optional[npt.ArrayLike] = None,
sensordata: Optional[npt.ArrayLike] = None,
chunk_size: Optional[int] = None,
persistent_pool: bool = False,
):
"""Rolls out open-loop trajectories from initial states, get subsequent states and sensor values.
@@ -44,8 +281,8 @@ def rollout(
Allocates outputs if none are given.
Args:
model: An mjModel or a sequence of MjModel with the same size signature.
data: An associated mjData instance.
model: An instance or length nroll sequence of MjModel with the same size signature.
data: Associated mjData instance or sequence of instances with length nthread.
initial_state: Array of initial states from which to roll out trajectories.
([nroll or 1] x nstate)
control: Open-loop controls array to apply during the rollouts.
@@ -59,6 +296,9 @@ def rollout(
(nroll x nstep x nstate)
sensordata: Sensor data output array (optional).
(nroll x nstep x nsensordata)
chunk_size: Determines threadpool chunk size. If unspecified,
chunk_size = max(1, nroll / (nthread * 10))
persistent_pool: Determines if a persistent thread pool is created or reused.
Returns:
state:
@@ -69,136 +309,41 @@ def rollout(
Raises:
ValueError: bad shapes or sizes.
""" # fmt: skip
# skip_checks shortcut:
# don't infer nroll/nstep
# don't support singleton expansion
# don't allocate output arrays
# just call rollout and return
if skip_checks:
_rollout.rollout(
if not isinstance(data, list):
data = [data] # Use a length 1 list to simplify code below
nthread = len(data) if len(data) > 1 else 0
# Use a persistent thread pool if requested
if persistent_pool:
# Create or restart persistent threadpool
global persistent_rollout
if persistent_rollout is None:
persistent_rollout = Rollout(nthread=nthread)
if persistent_rollout.nthread != nthread:
persistent_rollout.close()
persistent_rollout = Rollout(nthread=nthread)
rollout_ = persistent_rollout
else:
rollout_ = Rollout(nthread=nthread)
try:
return rollout_.rollout(
model,
data,
nstep,
control_spec,
initial_state,
initial_warmstart,
control,
state,
sensordata,
control_spec=control_spec,
skip_checks=skip_checks,
nstep=nstep,
initial_warmstart=initial_warmstart,
state=state,
sensordata=sensordata,
chunk_size=chunk_size,
)
return state, sensordata
if not isinstance(model, mujoco.MjModel):
model = list(model)
# check control_spec
if control_spec & ~mujoco.mjtState.mjSTATE_USER.value:
raise ValueError('control_spec can only contain bits in mjSTATE_USER')
# check types
if nstep and not isinstance(nstep, int):
raise ValueError('nstep must be an integer')
_check_must_be_numeric(
initial_state=initial_state,
initial_warmstart=initial_warmstart,
control=control,
state=state,
sensordata=sensordata,
)
# check number of dimensions
_check_number_of_dimensions(
2, initial_state=initial_state, initial_warmstart=initial_warmstart
)
_check_number_of_dimensions(
3, control=control, state=state, sensordata=sensordata
)
# ensure 2D, make contiguous, row-major (C ordering)
initial_state = _ensure_2d(initial_state)
initial_warmstart = _ensure_2d(initial_warmstart)
# ensure 3D, make contiguous, row-major (C ordering)
control = _ensure_3d(control)
state = _ensure_3d(state)
sensordata = _ensure_3d(sensordata)
# infer nroll, check for incompatibilities
nroll = _infer_dimension(
0,
1,
initial_state=initial_state,
initial_warmstart=initial_warmstart,
control=control,
state=state,
sensordata=sensordata,
)
if isinstance(model, list) and nroll == 1:
nroll = len(model)
if isinstance(model, list) and len(model) != nroll:
raise ValueError(
f'nroll inferred as {nroll} but model is length {len(model)}'
)
elif not isinstance(model, list):
model = [model] # Use a length 1 list to simplify code below
# infer nstep, check for incompatibilities
nstep = _infer_dimension(
1, nstep or 1, control=control, state=state, sensordata=sensordata
)
# get nstate/ncontrol/nv/nsensordata
# check that they are equal across models
nstate = mujoco.mj_stateSize(
model[0], mujoco.mjtState.mjSTATE_FULLPHYSICS.value
)
ncontrol = mujoco.mj_stateSize(model[0], control_spec)
nv = model[0].nv
nsensordata = model[0].nsensordata
for m in model[1:]:
if (
nstate
!= mujoco.mj_stateSize(m, mujoco.mjtState.mjSTATE_FULLPHYSICS.value)
or ncontrol != mujoco.mj_stateSize(m, control_spec)
or nv != m.nv
or nsensordata != m.nsensordata
):
raise ValueError('models are not compatible')
# check trailing dimensions
_check_trailing_dimension(nstate, initial_state=initial_state, state=state)
_check_trailing_dimension(ncontrol, control=control)
_check_trailing_dimension(nv, initial_warmstart=initial_warmstart)
_check_trailing_dimension(nsensordata, sensordata=sensordata)
# tile input arrays/lists if required (singleton expansion)
model = model * nroll if len(model) == 1 else model
initial_state = _tile_if_required(initial_state, nroll)
initial_warmstart = _tile_if_required(initial_warmstart, nroll)
control = _tile_if_required(control, nroll, nstep)
# allocate output if not provided
if state is None:
state = np.empty((nroll, nstep, nstate))
if sensordata is None:
sensordata = np.empty((nroll, nstep, nsensordata))
# call rollout
_rollout.rollout(
model,
data,
nstep,
control_spec,
initial_state,
initial_warmstart,
control,
state,
sensordata,
)
# return outputs
return state, sensordata
finally:
if not persistent_pool:
rollout_.close()
def _check_must_be_numeric(**kwargs):
+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):
+87
View File
@@ -0,0 +1,87 @@
// Copyright 2024 DeepMind Technologies Limited
//
// Licensed under the Apache License, Version 2.0 (the "License");
// you may not use this file except in compliance with the License.
// You may obtain a copy of the License at
//
// http://www.apache.org/licenses/LICENSE-2.0
//
// Unless required by applicable law or agreed to in writing, software
// distributed under the License is distributed on an "AS IS" BASIS,
// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
// See the License for the specific language governing permissions and
// limitations under the License.
#include "threadpool.h"
#include <condition_variable>
#include <functional>
#include <mutex>
#include <thread>
#include <utility>
#include <absl/base/attributes.h>
namespace mujoco::python {
ABSL_CONST_INIT thread_local int ThreadPool::worker_id_ = -1;
// ThreadPool constructor
ThreadPool::ThreadPool(int num_threads) : ctr_(0) {
for (int i = 0; i < num_threads; i++) {
threads_.push_back(std::thread(&ThreadPool::WorkerThread, this, i));
}
}
// ThreadPool destructor
ThreadPool::~ThreadPool() {
{
std::unique_lock<std::mutex> lock(m_);
for (int i = 0; i < threads_.size(); i++) {
queue_.push(nullptr);
}
cv_in_.notify_all();
}
for (auto& thread : threads_) {
thread.join();
}
}
// ThreadPool scheduler
void ThreadPool::Schedule(std::function<void()> task) {
std::unique_lock<std::mutex> lock(m_);
queue_.push(std::move(task));
cv_in_.notify_one();
}
// ThreadPool worker
void ThreadPool::WorkerThread(int i) {
worker_id_ = i;
while (true) {
auto task = [&]() {
std::unique_lock<std::mutex> lock(m_);
cv_in_.wait(lock, [&]() { return !queue_.empty(); });
std::function<void()> task = std::move(queue_.front());
queue_.pop();
cv_in_.notify_one();
return task;
}();
if (task == nullptr) {
{
std::unique_lock<std::mutex> lock(m_);
++ctr_;
cv_ext_.notify_one();
}
break;
}
task();
{
std::unique_lock<std::mutex> lock(m_);
++ctr_;
cv_ext_.notify_one();
}
}
}
} // namespace mujoco::python
+80
View File
@@ -0,0 +1,80 @@
// Copyright 2024 DeepMind Technologies Limited
//
// Licensed under the Apache License, Version 2.0 (the "License");
// you may not use this file except in compliance with the License.
// You may obtain a copy of the License at
//
// http://www.apache.org/licenses/LICENSE-2.0
//
// Unless required by applicable law or agreed to in writing, software
// distributed under the License is distributed on an "AS IS" BASIS,
// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
// See the License for the specific language governing permissions and
// limitations under the License.
#ifndef MUJOCO_PYTHON_THREADPOOL_H_
#define MUJOCO_PYTHON_THREADPOOL_H_
#include <condition_variable>
#include <cstdint>
#include <functional>
#include <mutex>
#include <queue>
#include <thread>
#include <vector>
#include <absl/base/attributes.h>
namespace mujoco::python {
// ThreadPool class
class ThreadPool {
public:
// constructor
explicit ThreadPool(int num_threads);
// destructor
~ThreadPool();
int NumThreads() const { return threads_.size(); }
// returns an ID between 0 and NumThreads() - 1. must be called within
// worker thread (returns -1 if not).
static int WorkerId() { return worker_id_; }
// ----- methods ----- //
// set task for threadpool
void Schedule(std::function<void()> task);
// return number of tasks completed
std::uint64_t GetCount() { return ctr_; }
// reset count to zero
void ResetCount() { ctr_ = 0; }
// wait for count, then return
void WaitCount(int value) {
std::unique_lock<std::mutex> lock(m_);
cv_ext_.wait(lock, [&]() { return this->GetCount() >= value; });
}
private:
// ----- methods ----- //
// execute task with available thread
void WorkerThread(int i);
ABSL_CONST_INIT static thread_local int worker_id_;
// ----- members ----- //
std::vector<std::thread> threads_;
std::mutex m_;
std::condition_variable cv_in_;
std::condition_variable cv_ext_;
std::queue<std::function<void()>> queue_;
std::uint64_t ctr_;
};
} // namespace mujoco::python
#endif // MUJOCO_PYTHON_THREADPOOL_H_