From a23a368778050afb1cf271ef5f28f4941ca8aad9 Mon Sep 17 00:00:00 2001 From: Yuval Tassa Date: Fri, 19 Jan 2024 09:02:27 -0800 Subject: [PATCH] clean up rollout.cc PiperOrigin-RevId: 599849967 Change-Id: I899b28bf78bf0414b07314c24cfe6e24441eba30 --- python/mujoco/rollout.cc | 32 +++++++++++++++----------------- 1 file changed, 15 insertions(+), 17 deletions(-) diff --git a/python/mujoco/rollout.cc b/python/mujoco/rollout.cc index 05d87ece..eacdc8fe 100644 --- a/python/mujoco/rollout.cc +++ b/python/mujoco/rollout.cc @@ -12,15 +12,14 @@ // See the License for the specific language governing permissions and // limitations under the License. -#include -#include #include #include #include -#include -#include "functions.h" +#include +#include "errors.h" #include "raw.h" +#include "structs.h" #include #include #include @@ -74,7 +73,6 @@ void _unsafe_rollout(const mjModel* m, mjData* d, int nstate, int nstep, // loop over initial states for (int s=0; s < nstate; s++) { - // set initial state if (state0) { mju_copy(d->qpos, state0 + s*nqva, nq); @@ -108,9 +106,9 @@ void _unsafe_rollout(const mjModel* m, mjData* d, int nstate, int nstep, mju_zero(d->xfrc_applied, 6*nbody); } if (!mocap) { - for (int j=0; jbody_mocapid[j]; - if (id>=0) { + if (id >= 0) { mju_copy3(d->mocap_pos+3*id, m->body_pos+3*j); mju_copy4(d->mocap_quat+4*id, m->body_quat+4*j); } @@ -174,7 +172,8 @@ mjtNum* get_array_ptr(std::optional> arg, int expected_size = nstate * nstep * dim; if (info.size != expected_size) { std::ostringstream msg; - msg << name << ".size should be " << expected_size << ", got " << info.size; + msg << name << ".size should be " << expected_size << + ", got " << info.size; throw py::value_error(msg.str()); } return static_cast(info.ptr); @@ -200,7 +199,6 @@ PYBIND11_MODULE(_rollout, pymodule) { std::optional state, std::optional sensordata ) { - const raw::MjModel* model = m.get(); raw::MjData* data = d.get(); @@ -213,7 +211,8 @@ PYBIND11_MODULE(_rollout, pymodule) { int nqva = model->nq + model->nv + model->na; mjtNum* init_state_ptr = get_array_ptr(init_state, "initial_state", nstate, 1, nqva); - mjtNum* ctrl_ptr = get_array_ptr(ctrl, "ctrl", nstate, nstep, model->nu); + mjtNum* ctrl_ptr = + get_array_ptr(ctrl, "ctrl", nstate, nstep, model->nu); mjtNum* qfrc_ptr = get_array_ptr(qfrc, "qfrc_applied", nstate, nstep, model->nv); mjtNum* xfrc_ptr = @@ -222,11 +221,11 @@ PYBIND11_MODULE(_rollout, pymodule) { get_array_ptr(mocap, "mocap", nstate, nstep, 7*model->nmocap); mjtNum* init_time_ptr = get_array_ptr(init_time, "init_time", nstate, 1, 1); - mjtNum* init_warmstart_ptr = - get_array_ptr(init_warmstart, "init_warmstart", nstate, 1, model->nv); + mjtNum* init_warmstart_ptr = get_array_ptr( + init_warmstart, "init_warmstart", nstate, 1, model->nv); mjtNum* state_ptr = get_array_ptr(state, "state", nstate, nstep, nqva); - mjtNum* sensordata_ptr = - get_array_ptr(sensordata, "sensordata", nstate, nstep, model->nsensordata); + mjtNum* sensordata_ptr = get_array_ptr(sensordata, "sensordata", nstate, + nstep, model->nsensordata); // perform rollouts { @@ -255,9 +254,8 @@ PYBIND11_MODULE(_rollout, pymodule) { py::arg("sensordata") = py::none(), py::doc(rollout_doc) ); +} } // namespace -} - -} +} // namespace mujoco::python