rollout accepts a list of models of length nroll

This commit is contained in:
Levi Burner
2024-11-24 16:50:28 -05:00
parent 6e129a5ab2
commit 943eb6bc7e
5 changed files with 102 additions and 43 deletions
+29 -26
View File
@@ -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);
}
},