diff --git a/python/mujoco/experimental/studio/sim.cc b/python/mujoco/experimental/studio/sim.cc index ae75612b..c694abaa 100644 --- a/python/mujoco/experimental/studio/sim.cc +++ b/python/mujoco/experimental/studio/sim.cc @@ -42,11 +42,21 @@ PYBIND11_MODULE(sim, m) { .def(py::init<>()) .def( "advance", - [](StepControl& self, mujoco::python::MjModelWrapper& model, - mujoco::python::MjDataWrapper& data) { - return self.Advance(model.get(), data.get()); + [](StepControl& self, py::object model_obj, py::object data_obj, + py::object step_fn) { + auto& model = py::cast(model_obj); + auto& data = py::cast(data_obj); + if (step_fn.is_none()) { + return self.Advance(model.get(), data.get()); + } else { + return self.Advance( + model.get(), data.get(), + [step_fn, model_obj, data_obj](mjModel*, mjData*) { + step_fn(model_obj, data_obj); + }); + } }, - py::arg("model"), py::arg("data"), + py::arg("model"), py::arg("data"), py::arg("step_fn") = py::none(), "Step physics forward, respecting speed settings and refresh budget.") .def("force_sync", &StepControl::ForceSync, "Ensures the next Advance() will synchronize time and step once.") diff --git a/src/experimental/platform/sim/step_control.cc b/src/experimental/platform/sim/step_control.cc index 1f5b54c3..a41a6bbe 100644 --- a/src/experimental/platform/sim/step_control.cc +++ b/src/experimental/platform/sim/step_control.cc @@ -17,6 +17,7 @@ #include #include #include +#include #include #include @@ -91,7 +92,8 @@ StepControl::PauseState StepControl::GetPauseState() const { return pause_state_; } -StepControl::Status StepControl::Advance(mjModel* m, mjData* d) { +StepControl::Status StepControl::Advance(mjModel* m, mjData* d, + StepFn step_fn) { if (!m) { return Status::kOk; } @@ -182,7 +184,11 @@ StepControl::Status StepControl::Advance(mjModel* m, mjData* d) { mjtNum prev_time = d->time; InjectNoise(m, d); - mj_step(m, d); + if (step_fn) { + step_fn(m, d); + } else { + mj_step(m, d); + } if (mjDISABLED(mjDSBL_AUTORESET)) { for (mjtWarning w : kDivergedWarnings) { diff --git a/src/experimental/platform/sim/step_control.h b/src/experimental/platform/sim/step_control.h index 1f34a810..1717c975 100644 --- a/src/experimental/platform/sim/step_control.h +++ b/src/experimental/platform/sim/step_control.h @@ -16,6 +16,7 @@ #define MUJOCO_SRC_EXPERIMENTAL_PLATFORM_SIM_STEP_CONTROL_H_ #include +#include #include #include @@ -24,6 +25,7 @@ namespace mujoco::platform { using Seconds = std::chrono::duration; using Clock = std::chrono::steady_clock; +using StepFn = std::function; // State and logic for physics synchronization and stepping. class StepControl { @@ -53,7 +55,7 @@ class StepControl { mjWARN_BADQACC, mjWARN_BADQVEL, mjWARN_BADQPOS}; // Steps physics forward, respecting speed settings and refresh budget. - Status Advance(mjModel* m, mjData* d); + Status Advance(mjModel* m, mjData* d, StepFn step_fn = nullptr); // Ensures the next call to Advance() will synchronize time and step once. void ForceSync();