// Copyright 2022 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 #include // NOLINT(build/c++11) #include #include #include #include #include // NOLINT(build/c++11) #include #include #include #include #include #include #include #include "errors.h" #include "indexers.h" #include "structs.h" #include #include #include #include namespace mujoco::python { namespace { using UIAdapter = mujoco::GlfwAdapter; namespace py = ::pybind11; template constexpr inline std::size_t sizeof_arr(const T (&arr)[N]) { return sizeof(arr); } template class UIAdapterWithPyCallback : public Adapter { public: template UIAdapterWithPyCallback(py::handle key_callback, Args&&... args) : Adapter(std::forward(args)...) { if (!key_callback.is_none()) { Py_XINCREF(key_callback.ptr()); key_callback_ = key_callback.ptr(); } } ~UIAdapterWithPyCallback() override { Py_XDECREF(key_callback_); } protected: void OnKey(int key, int scancode, int act) override { Adapter::OnKey(key, scancode, act); if (this->IsKeyDownEvent(act) && key_callback_) { py::gil_scoped_acquire gil; (py::handle(key_callback_))(this->last_key_); } } private: PyObject* key_callback_ = nullptr; }; class SimulateWrapper { public: SimulateWrapper(std::unique_ptr platform_ui_adapter, py::object cam, py::object opt, py::object pert, py::object user_scn, bool is_passive) : simulate_(new mujoco::Simulate( std::move(platform_ui_adapter), cam.cast().get(), opt.cast().get(), pert.cast().get(), is_passive)), m_(py::none()), d_(py::none()), cam_(cam), opt_(opt), pert_(pert), user_scn_(user_scn) { if (!user_scn.is_none()) { simulate_->user_scn = user_scn_.cast().get(); } } ~SimulateWrapper() { Destroy(); } void Destroy() { if (simulate_) { ClearImages(); delete simulate_; simulate_ = nullptr; destroyed_.store(1); } } void WaitUntilExit() { // TODO: replace with atomic wait when we migrate to C++20 while (simulate_ && simulate_->exitrequest.load() != 2) { std::this_thread::sleep_for(std::chrono::milliseconds(10)); } } void Load(py::object m, py::object d, const std::string& path) { if (!simulate_) { return; } mjModel* m_raw = m.cast().get(); mjData* d_raw = d.cast().get(); { py::gil_scoped_release no_gil; simulate_->Load(m_raw, d_raw, path.c_str()); } m_ = m; d_ = d; m_raw_ = m_raw; d_raw_ = d_raw; } mujoco::Simulate* simulate() { return simulate_; } py::object GetModel() const { return m_; } py::object GetData() const { return d_; } mjrRect GetViewport() const { // Return the viewport corresponding to the 3D view, i.e. the viewer window // without the UI elements. return simulate_->uistate.rect[3]; } void SetFigures( const std::vector>& viewports_figures) { // Pairs of [viewport, figure], where viewport corresponds to the location // of the figure on the viewer window. std::vector> user_figures; for (const auto& [viewport, figure] : viewports_figures) { mjvFigure casted_figure = *figure.cast().get(); user_figures.push_back(std::make_pair(viewport, casted_figure)); } // Set them all at once to prevent figure flickering. simulate_->user_figures_ = user_figures; } void ClearFigures() { simulate_->user_figures_.clear(); } void SetTexts( const std::vector>& texts) { // Collection of [font, gridpos, text1, text2] tuples for overlay text std::vector> user_texts; for (const auto& [font, gridpos, text1, text2] : texts) { user_texts.push_back(std::make_tuple(font, gridpos, text1, text2)); } // Set them all at once to prevent text flickering. simulate_->user_texts_ = user_texts; } void ClearTexts() { simulate_->user_texts_.clear(); } void SetImages( const std::vector> viewports_images ) { // Clear previous images to prevent memory leaks ClearImages(); for (const auto& [viewport, image] : viewports_images) { auto buf = image.request(); if (buf.ndim != 3) { throw std::invalid_argument("image must have 3 dimensions (H, W, C)"); } if (static_cast(buf.shape[2]) != 3) { throw std::invalid_argument("image must have 3 channels"); } if (buf.itemsize != sizeof(unsigned char)) { throw std::invalid_argument("image must be uint8 format"); } // Calculate size of the image data size_t height = buf.shape[0]; size_t width = buf.shape[1]; size_t size = height * width * 3; // Make a copy of the image data to prevent flickering unsigned char* image_copy = new unsigned char[size]; std::memcpy(image_copy, buf.ptr, size); simulate_->user_images_.push_back(std::make_tuple(viewport, image_copy)); } } void ClearImages() { // Free memory for each image before clearing the vector for (const auto& [viewport, image_ptr] : simulate_->user_images_) { delete[] image_ptr; } simulate_->user_images_.clear(); } private: mujoco::Simulate* simulate_; std::atomic_int destroyed_ = 0; // Hold references to keep these Python objects alive for as long as the // simulate object. py::object m_; py::object d_; py::object cam_; py::object opt_; py::object pert_; py::object user_scn_; mjModel* m_raw_ = nullptr; mjData* d_raw_ = nullptr; }; inline mujoco::Simulate& SimulateRefOrThrow(SimulateWrapper& wrapper) { auto* sim = wrapper.simulate(); if (!sim) { throw UnexpectedError("simulate object is already deleted"); } return *sim; } template inline auto CallIfNotNull(T (*func)(mujoco::Simulate&, Args...)) { return [func](SimulateWrapper& wrapper, Args&&... args) { return func(SimulateRefOrThrow(wrapper), std::forward(args)...); }; } template inline auto CallIfNotNull(void (*func)(mujoco::Simulate&, Args...)) { return [func](SimulateWrapper& wrapper, Args&&... args) -> void { func(SimulateRefOrThrow(wrapper), std::forward(args)...); }; } template inline auto CallIfNotNull(void (mujoco::Simulate::*func)(Args...)) { return [func](SimulateWrapper& wrapper, Args&&... args) -> void { (SimulateRefOrThrow(wrapper).*func)(std::forward(args)...); }; } template inline auto GetIfNotNull(T mujoco::Simulate::* member) { return [member](SimulateWrapper& wrapper) -> T& { return SimulateRefOrThrow(wrapper).*member; }; } template inline auto SetIfNotNull(T mujoco::Simulate::* member) { return [member](SimulateWrapper& wrapper, const T& value) -> void { SimulateRefOrThrow(wrapper).*member = value; }; } PYBIND11_MODULE(_simulate, pymodule) { py::class_(pymodule, "Mutex") .def( "__enter__", [](SimulateMutex& mtx) { mtx.lock(); }, py::call_guard()) .def( "__exit__", [](SimulateMutex& mtx, py::handle, py::handle, py::handle) { mtx.unlock(); }, py::call_guard()); py::class_(pymodule, "Simulate") .def_readonly_static("MAX_GEOM", &mujoco::Simulate::kMaxGeom) .def(py::init([](py::object scn, py::object cam, py::object opt, py::object pert, bool run_physics_thread, py::object key_callback) { bool is_passive = !run_physics_thread; return std::make_unique( std::make_unique>(key_callback), scn, cam, opt, pert, is_passive); })) .def("destroy", &SimulateWrapper::Destroy) .def("load_message", CallIfNotNull(&mujoco::Simulate::LoadMessage), py::call_guard()) .def("load", &SimulateWrapper::Load) .def("load_message_clear", CallIfNotNull(&mujoco::Simulate::LoadMessageClear), py::call_guard()) .def("sync", CallIfNotNull(&mujoco::Simulate::Sync), py::call_guard()) .def("add_to_history", CallIfNotNull(&mujoco::Simulate::AddToHistory), py::call_guard()) .def("render_loop", CallIfNotNull(&mujoco::Simulate::RenderLoop), py::call_guard()) .def("lock", GetIfNotNull(&mujoco::Simulate::mtx), py::call_guard(), py::return_value_policy::reference_internal) .def("set_figures", &SimulateWrapper::SetFigures, py::arg("viewports_figures")) .def("clear_figures", &SimulateWrapper::ClearFigures) .def("set_texts", &SimulateWrapper::SetTexts, py::arg("overlay_texts")) .def("clear_texts", &SimulateWrapper::ClearTexts) .def("set_images", &SimulateWrapper::SetImages, py::arg("viewports_images")) .def("clear_images", &SimulateWrapper::ClearImages) .def_property_readonly("m", &SimulateWrapper::GetModel) .def_property_readonly("d", &SimulateWrapper::GetData) .def_property_readonly("viewport", &SimulateWrapper::GetViewport) .def_property_readonly("ctrl_noise_std", GetIfNotNull(&mujoco::Simulate::ctrl_noise_std), py::call_guard()) .def_property_readonly("ctrl_noise_rate", GetIfNotNull(&mujoco::Simulate::ctrl_noise_rate), py::call_guard()) .def_property_readonly("real_time_index", GetIfNotNull(&mujoco::Simulate::real_time_index), py::call_guard()) .def_property("speed_changed", GetIfNotNull(&mujoco::Simulate::speed_changed), SetIfNotNull(&mujoco::Simulate::speed_changed), py::call_guard()) .def_property("measured_slowdown", GetIfNotNull(&mujoco::Simulate::measured_slowdown), SetIfNotNull(&mujoco::Simulate::measured_slowdown), py::call_guard()) .def_property_readonly("refresh_rate", GetIfNotNull(&mujoco::Simulate::refresh_rate), py::call_guard()) .def_property_readonly("busywait", GetIfNotNull(&mujoco::Simulate::busywait), py::call_guard()) .def_property_readonly("run", GetIfNotNull(&mujoco::Simulate::run), py::call_guard()) .def_property_readonly("exitrequest", CallIfNotNull(+[](mujoco::Simulate& sim) { return sim.exitrequest.load(); }), py::call_guard()) .def("exit", [](SimulateWrapper& wrapper) { mujoco::Simulate* sim = wrapper.simulate(); if (!sim) { return; } int value = 0; sim->exitrequest.compare_exchange_strong(value, 1); wrapper.WaitUntilExit(); }) .def_property_readonly("uiloadrequest", CallIfNotNull(+[](mujoco::Simulate& sim) { return sim.uiloadrequest.load(); }), py::call_guard()) .def("uiloadrequest_decrement", CallIfNotNull(+[](mujoco::Simulate& sim) { sim.uiloadrequest.fetch_sub(1); }), py::call_guard()) .def("update_hfield", CallIfNotNull(+[](mujoco::Simulate& sim, int hfieldid) { sim.UpdateHField(hfieldid); }), py::call_guard()) .def("update_mesh", CallIfNotNull(+[](mujoco::Simulate& sim, int meshid) { sim.UpdateMesh(meshid); }), py::call_guard()) .def("update_texture", CallIfNotNull(+[](mujoco::Simulate& sim, int texid) { sim.UpdateTexture(texid); }), py::call_guard()) .def_property( "droploadrequest", CallIfNotNull(+[](mujoco::Simulate& sim) { return sim.droploadrequest.load(); }), CallIfNotNull(+[](mujoco::Simulate& sim, bool droploadrequest) { sim.droploadrequest.store(droploadrequest); }), py::call_guard()) .def_property_readonly("dropfilename", GetIfNotNull(&mujoco::Simulate::dropfilename), py::call_guard()) .def_property_readonly("filename", GetIfNotNull(&mujoco::Simulate::filename), py::call_guard()) .def_property( "load_error", GetIfNotNull(&mujoco::Simulate::load_error), CallIfNotNull(+[](mujoco::Simulate& sim, const std::string& error) { const auto max_length = sizeof_arr(sim.load_error); std::strncpy(sim.load_error, error.c_str(), max_length - 1); sim.load_error[max_length - 1] = '\0'; })) .def_property("ui0_enable", GetIfNotNull(&mujoco::Simulate::ui0_enable), CallIfNotNull(+[](mujoco::Simulate& sim, int enabled) { sim.ui0_enable = enabled; }), py::call_guard()) .def_property("ui1_enable", GetIfNotNull(&mujoco::Simulate::ui1_enable), CallIfNotNull(+[](mujoco::Simulate& sim, int enabled) { sim.ui1_enable = enabled; }), py::call_guard()); pymodule.def("set_glfw_dlhandle", [](std::uintptr_t dlhandle) { mujoco::Glfw(reinterpret_cast(dlhandle)); }); } } // namespace } // namespace mujoco::python