rollout accepts a list of models of length nroll
This commit is contained in:
+29
-26
@@ -37,7 +37,7 @@ const auto rollout_doc = R"(
|
||||
Roll out open-loop trajectories from initial states, get resulting states and sensor values.
|
||||
|
||||
input arguments (required):
|
||||
model instance of MjModel
|
||||
model list of MjModel instances of length nroll
|
||||
data associated instance of MjData
|
||||
nstep integer, number of steps to be taken for each trajectory
|
||||
control_spec specification of controls, ncontrol = mj_stateSize(m, control_spec)
|
||||
@@ -54,18 +54,18 @@ Roll out open-loop trajectories from initial states, get resulting states and se
|
||||
// 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(const mjModel* m, mjData* d, int nroll, int nstep, unsigned int control_spec,
|
||||
void _unsafe_rollout(const mjModel** m, mjData* d, int nroll, int nstep, unsigned int control_spec,
|
||||
const mjtNum* state0, const mjtNum* warmstart0, const mjtNum* control,
|
||||
mjtNum* state, mjtNum* sensordata) {
|
||||
// sizes
|
||||
int nstate = mj_stateSize(m, mjSTATE_FULLPHYSICS);
|
||||
int ncontrol = mj_stateSize(m, control_spec);
|
||||
int nv = m->nv, nbody = m->nbody, neq = m->neq;
|
||||
int nsensordata = m->nsensordata;
|
||||
int nstate = mj_stateSize(m[0], mjSTATE_FULLPHYSICS);
|
||||
int ncontrol = mj_stateSize(m[0], control_spec);
|
||||
int nv = m[0]->nv, nbody = m[0]->nbody, neq = m[0]->neq;
|
||||
int nsensordata = m[0]->nsensordata;
|
||||
|
||||
// clear user inputs if unspecified
|
||||
if (!(control_spec & mjSTATE_CTRL)) {
|
||||
mju_zero(d->ctrl, m->nu);
|
||||
mju_zero(d->ctrl, m[0]->nu);
|
||||
}
|
||||
if (!(control_spec & mjSTATE_QFRC_APPLIED)) {
|
||||
mju_zero(d->qfrc_applied, nv);
|
||||
@@ -75,26 +75,26 @@ void _unsafe_rollout(const mjModel* m, mjData* d, int nroll, int nstep, unsigned
|
||||
}
|
||||
if (!(control_spec & mjSTATE_MOCAP_POS)) {
|
||||
for (int i = 0; i < nbody; i++) {
|
||||
int id = m->body_mocapid[i];
|
||||
if (id >= 0) mju_copy3(d->mocap_pos+3*id, m->body_pos+3*i);
|
||||
int id = m[0]->body_mocapid[i];
|
||||
if (id >= 0) mju_copy3(d->mocap_pos+3*id, m[0]->body_pos+3*i);
|
||||
}
|
||||
}
|
||||
if (!(control_spec & mjSTATE_MOCAP_QUAT)) {
|
||||
for (int i = 0; i < nbody; i++) {
|
||||
int id = m->body_mocapid[i];
|
||||
if (id >= 0) mju_copy4(d->mocap_quat+4*id, m->body_quat+4*i);
|
||||
int id = m[0]->body_mocapid[i];
|
||||
if (id >= 0) mju_copy4(d->mocap_quat+4*id, m[0]->body_quat+4*i);
|
||||
}
|
||||
}
|
||||
if (!(control_spec & mjSTATE_EQ_ACTIVE)) {
|
||||
for (int i = 0; i < neq; i++) {
|
||||
d->eq_active[i] = m->eq_active0[i];
|
||||
d->eq_active[i] = m[0]->eq_active0[i];
|
||||
}
|
||||
}
|
||||
|
||||
// loop over rollouts
|
||||
for (int r = 0; r < nroll; r++) {
|
||||
// set initial state
|
||||
mj_setState(m, d, state0 + r*nstate, mjSTATE_FULLPHYSICS);
|
||||
mj_setState(m[r], d, state0 + r*nstate, mjSTATE_FULLPHYSICS);
|
||||
|
||||
// set warmstart accelerations
|
||||
if (warmstart0) {
|
||||
@@ -124,7 +124,7 @@ void _unsafe_rollout(const mjModel* m, mjData* d, int nroll, int nstep, unsigned
|
||||
for (; t < nstep; t++) {
|
||||
int step = r*nstep + t;
|
||||
if (state) {
|
||||
mj_getState(m, d, state + step*nstate, mjSTATE_FULLPHYSICS);
|
||||
mj_getState(m[r], d, state + step*nstate, mjSTATE_FULLPHYSICS);
|
||||
}
|
||||
if (sensordata) {
|
||||
mju_copy(sensordata + step*nsensordata, d->sensordata, nsensordata);
|
||||
@@ -137,15 +137,15 @@ void _unsafe_rollout(const mjModel* m, mjData* d, int nroll, int nstep, unsigned
|
||||
|
||||
// controls
|
||||
if (control) {
|
||||
mj_setState(m, d, control + step*ncontrol, control_spec);
|
||||
mj_setState(m[r], d, control + step*ncontrol, control_spec);
|
||||
}
|
||||
|
||||
// step
|
||||
mj_step(m, d);
|
||||
mj_step(m[r], d);
|
||||
|
||||
// copy out new state
|
||||
if (state) {
|
||||
mj_getState(m, d, state + step*nstate, mjSTATE_FULLPHYSICS);
|
||||
mj_getState(m[r], d, state + step*nstate, mjSTATE_FULLPHYSICS);
|
||||
}
|
||||
|
||||
// copy out sensor values
|
||||
@@ -188,7 +188,7 @@ PYBIND11_MODULE(_rollout, pymodule) {
|
||||
// get subsequent states and corresponding sensor values
|
||||
pymodule.def(
|
||||
"rollout",
|
||||
[](const MjModelWrapper& m, MjDataWrapper& d,
|
||||
[](py::list m, MjDataWrapper& d,
|
||||
int nstep, unsigned int control_spec,
|
||||
const PyCArray state0,
|
||||
std::optional<const PyCArray> warmstart0,
|
||||
@@ -196,7 +196,12 @@ PYBIND11_MODULE(_rollout, pymodule) {
|
||||
std::optional<const PyCArray> state,
|
||||
std::optional<const PyCArray> sensordata
|
||||
) {
|
||||
const raw::MjModel* model = m.get();
|
||||
// get raw pointers
|
||||
int nroll = state0.shape(0);
|
||||
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
|
||||
@@ -205,19 +210,17 @@ PYBIND11_MODULE(_rollout, pymodule) {
|
||||
}
|
||||
|
||||
// get sizes
|
||||
int nstate = mj_stateSize(model, mjSTATE_FULLPHYSICS);
|
||||
int ncontrol = mj_stateSize(model, control_spec);
|
||||
int nroll = state0.shape(0);
|
||||
int nstate = mj_stateSize(model_ptrs[0], mjSTATE_FULLPHYSICS);
|
||||
int ncontrol = mj_stateSize(model_ptrs[0], control_spec);
|
||||
|
||||
// get raw pointers
|
||||
mjtNum* state0_ptr = get_array_ptr(state0, "state0", nroll, 1, nstate);
|
||||
mjtNum* warmstart0_ptr = get_array_ptr(warmstart0, "warmstart0", nroll,
|
||||
1, model->nv);
|
||||
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->nsensordata);
|
||||
nstep, model_ptrs[0]->nsensordata);
|
||||
|
||||
// perform rollouts
|
||||
{
|
||||
@@ -226,7 +229,7 @@ PYBIND11_MODULE(_rollout, pymodule) {
|
||||
|
||||
// call unsafe rollout function
|
||||
InterceptMjErrors(_unsafe_rollout)(
|
||||
model, data, nroll, nstep, control_spec, state0_ptr,
|
||||
model_ptrs, data, nroll, nstep, control_spec, state0_ptr,
|
||||
warmstart0_ptr, control_ptr, state_ptr, sensordata_ptr);
|
||||
}
|
||||
},
|
||||
|
||||
Reference in New Issue
Block a user