From eadd13038a144a806841ce0c00b32f578ee825a4 Mon Sep 17 00:00:00 2001 From: Sam Haves Date: Wed, 9 Apr 2025 09:27:40 -0700 Subject: [PATCH] Move Mujoco struct wrappers definitions into library separate from pybind module. This allows other modules to depend on functions such as FromRawPointer that previously were only defined in the _struct python module. PiperOrigin-RevId: 745619670 Change-Id: I9f1af60cb665c00d853b7ca38201b37302d278af --- python/mujoco/CMakeLists.txt | 20 + python/mujoco/structs.cc | 1616 +++-------------------------- python/mujoco/structs_wrappers.cc | 1328 ++++++++++++++++++++++++ 3 files changed, 1505 insertions(+), 1459 deletions(-) create mode 100644 python/mujoco/structs_wrappers.cc diff --git a/python/mujoco/CMakeLists.txt b/python/mujoco/CMakeLists.txt index 2a59248e..ac15ccd5 100644 --- a/python/mujoco/CMakeLists.txt +++ b/python/mujoco/CMakeLists.txt @@ -278,6 +278,25 @@ target_link_libraries( raw ) +add_library(structs_wrappers STATIC + structs_wrappers.cc + serialization.h + indexers.cc +) +target_include_directories(structs_wrappers PRIVATE ${Python3_INCLUDE_DIRS}) +target_link_libraries( + structs_wrappers + PRIVATE absl::flat_hash_map + absl::span + crossplatform + mujoco + raw + structs_header + pybind11::headers + Eigen3::Eigen +) + + add_library(functions_header INTERFACE) target_sources(functions_header INTERFACE functions.h) set_target_properties(functions_header PROPERTIES PUBLIC_HEADER functions.h) @@ -403,6 +422,7 @@ target_link_libraries( func_wrap function_traits structs_header + structs_wrappers ) if(NOT EXISTS ${CMAKE_CURRENT_SOURCE_DIR}/specs.cc.inc) diff --git a/python/mujoco/structs.cc b/python/mujoco/structs.cc index 77200386..3cf6eb82 100644 --- a/python/mujoco/structs.cc +++ b/python/mujoco/structs.cc @@ -16,37 +16,23 @@ #include -#include -#include -#include // NOLINT(build/c++11) -#include -#include #include -#include -#include #include #include #include -#include #include #include #include -#include -#include -#include #include #include -#include #include #include #include "errors.h" #include "function_traits.h" #include "indexer_xmacro.h" #include "indexers.h" -#include "private.h" #include "raw.h" -#include "serialization.h" #include #include #include @@ -69,13 +55,9 @@ namespace { // (dim0, dim1). #define X_ARRAY_SHAPE(dim0, dim1) XArrayShapeImpl(#dim1)((dim0), (dim1)) -std::vector XArrayShapeImpl1D(int dim0, int dim1) { - return {dim0}; -} +std::vector XArrayShapeImpl1D(int dim0, int dim1) { return {dim0}; } -std::vector XArrayShapeImpl2D(int dim0, int dim1) { - return {dim0, dim1}; -} +std::vector XArrayShapeImpl2D(int dim0, int dim1) { return {dim0, dim1}; } constexpr auto XArrayShapeImpl(const std::string_view dim1_str) { if (dim1_str == "1") { @@ -85,334 +67,6 @@ constexpr auto XArrayShapeImpl(const std::string_view dim1_str) { } } -inline std::size_t NConMax(const mjData* d) { - return d->narena / sizeof(mjContact); -} - -} // namespace - -// ==================== MJOPTION =============================================== -#define X(var, dim) , var(InitPyArray(std::array{dim}, ptr_->var, owner_)) -MjOptionWrapper::MjWrapper() - : WrapperBase([]() { - raw::MjOption* const opt = new raw::MjOption; - mj_defaultOption(opt); - return opt; - }()) - MJOPTION_VECTORS {} - -MjOptionWrapper::MjWrapper(raw::MjOption* ptr, py::handle owner) - : WrapperBase(ptr, owner) - MJOPTION_VECTORS {} -#undef X - -MjOptionWrapper::MjWrapper(const MjOptionWrapper& other) - : MjOptionWrapper() { - *this->ptr_ = *other.ptr_; -} - -// ==================== MJVISUAL =============================================== -#define X(var) var(InitPyArray(ptr_->var, owner_)) -MjVisualHeadlightWrapper::MjWrapper() - : WrapperBase(new raw::MjVisualHeadlight{}), - X(ambient), - X(diffuse), - X(specular) {} - -MjVisualHeadlightWrapper::MjWrapper( - raw::MjVisualHeadlight* ptr, py::handle owner) - : WrapperBase(ptr, owner), - X(ambient), - X(diffuse), - X(specular) {} -#undef X - -MjVisualHeadlightWrapper::MjWrapper(const MjVisualHeadlightWrapper& other) - : MjVisualHeadlightWrapper() { - *this->ptr_ = *other.ptr_; -} - -#define X(var) var(InitPyArray(ptr_->var, owner_)) -MjVisualRgbaWrapper::MjWrapper() - : WrapperBase(new raw::MjVisualRgba{}), - X(fog), - X(haze), - X(force), - X(inertia), - X(joint), - X(actuator), - X(actuatornegative), - X(actuatorpositive), - X(com), - X(camera), - X(light), - X(selectpoint), - X(connect), - X(contactpoint), - X(contactforce), - X(contactfriction), - X(contacttorque), - X(contactgap), - X(rangefinder), - X(constraint), - X(slidercrank), - X(crankbroken), - X(frustum) {} - -MjVisualRgbaWrapper::MjWrapper(raw::MjVisualRgba* ptr, py::handle owner) - : WrapperBase(ptr, owner), - X(fog), - X(haze), - X(force), - X(inertia), - X(joint), - X(actuator), - X(actuatornegative), - X(actuatorpositive), - X(com), - X(camera), - X(light), - X(selectpoint), - X(connect), - X(contactpoint), - X(contactforce), - X(contactfriction), - X(contacttorque), - X(contactgap), - X(rangefinder), - X(constraint), - X(slidercrank), - X(crankbroken), - X(frustum) {} -#undef X - -MjVisualRgbaWrapper::MjWrapper(const MjVisualRgbaWrapper& other) - : MjVisualRgbaWrapper() { - *this->ptr_ = *other.ptr_; -} - -MjVisualWrapper::MjWrapper() - : WrapperBase(new raw::MjVisual{}), - headlight(&ptr_->headlight, owner_), - rgba(&ptr_->rgba, owner_) {} - -MjVisualWrapper::MjWrapper(raw::MjVisual* ptr, py::handle owner) - : WrapperBase(ptr, owner), - headlight(&ptr_->headlight, owner_), - rgba(&ptr_->rgba, owner_) {} - - -MjVisualWrapper::MjWrapper(const MjVisualWrapper& other) - : MjVisualWrapper() { - *this->ptr_ = *other.ptr_; -} - -// ==================== MJMODEL ================================================ -static void MjModelCapsuleDestructor(PyObject* pyobj) { - mj_deleteModel( - static_cast(PyCapsule_GetPointer(pyobj, nullptr))); -} - -static absl::flat_hash_map& -MjModelRawPointerMap() { - static auto* hash_map = - new absl::flat_hash_map(); - return *hash_map; -} - -MjModelWrapper* MjModelWrapper::FromRawPointer(raw::MjModel* m) noexcept { - try { - auto& map = MjModelRawPointerMap(); - { - py::gil_scoped_acquire gil; - auto found = map.find(m); - return found != map.end() ? found->second : nullptr; - } - } catch (...) { - return nullptr; - } -} - -#undef MJ_M -#define MJ_M(x) ptr_->x -#define X(dtype, var, dim0, dim1) \ - , var (InitPyArray(X_ARRAY_SHAPE(ptr_->dim0, dim1), ptr_->var, owner_)) -MjModelWrapper::MjWrapper(raw::MjModel* ptr) - : WrapperBase(ptr, &MjModelCapsuleDestructor), - opt(&ptr->opt, owner_), - vis(&ptr->vis, owner_), - stat(&ptr->stat, owner_) - MJMODEL_POINTERS, - text_data_bytes(ptr->text_data, ptr->ntextdata), - names_bytes(ptr->names, ptr->nnames), - paths_bytes(ptr->paths, ptr->npaths), - indexer_(ptr, owner_) { - bool is_newly_inserted = false; - { - py::gil_scoped_acquire gil; - is_newly_inserted = MjModelRawPointerMap().insert({ptr_, this}).second; - } - if (!is_newly_inserted) { - throw UnexpectedError( - "MjModelWrapper(mjModel*): MjModelRawPointerMap already contains this " - "raw mjModel*"); - } -} - -MjModelWrapper::MjWrapper(MjModelWrapper&& other) - : WrapperBase(other.ptr_, other.owner_), - opt(&ptr_->opt, owner_), - vis(&ptr_->vis, owner_), - stat(&ptr_->stat, owner_) - MJMODEL_POINTERS, - text_data_bytes(ptr_->text_data, ptr_->ntextdata), - names_bytes(ptr_->names, ptr_->nnames), - paths_bytes(ptr_->paths, ptr_->npaths), - indexer_(ptr_, owner_) { - bool is_newly_inserted = false; - { - py::gil_scoped_acquire gil; - is_newly_inserted = - MjModelRawPointerMap().insert_or_assign(ptr_, this).second; - } - if (is_newly_inserted) { - throw UnexpectedError( - "MjModelRawPointerMap does not contains the moved-from mjModel*"); - } - other.ptr_ = nullptr; -} -#undef X -#undef MJ_M -#define MJ_M(x) x - -// Delegating to the MjModelWrapper::MjWrapper(raw::MjModel*) constructor, -// no need to modify MjModelRawPointerMap here. -MjModelWrapper::MjWrapper(const MjModelWrapper& other) - : MjModelWrapper(InterceptMjErrors(mj_copyModel)(NULL, other.get())) {} - -MjModelWrapper::~MjWrapper() { - if (ptr_) { - bool erased = false; - { - py::gil_scoped_acquire gil; - erased = MjModelRawPointerMap().erase(ptr_); - } - if (!erased) { - std::cerr << "MjModelRawPointerMap does not contain this raw mjModel*" - << std::endl; - std::terminate(); - } - } -} - -// Helper function for both LoadXMLFile and LoadBinaryFile. -// Creates a temporary MJB from the assets dictionary if one is supplied. -template -static raw::MjModel* LoadModelFileImpl( - const std::string& filename, - const std::vector& assets, - LoadFunc&& loadfunc) { - mjVFS vfs; - mjVFS* vfs_ptr = nullptr; - if (!assets.empty()) { - mj_defaultVFS(&vfs); - vfs_ptr = &vfs; - for (const auto& asset : assets) { - std::string buffer_name = StripPath(asset.name); - const int vfs_error = InterceptMjErrors(mj_addBufferVFS)( - vfs_ptr, buffer_name.c_str(), asset.content, asset.content_size); - if (vfs_error) { - mj_deleteVFS(vfs_ptr); - if (vfs_error == 2) { - throw py::value_error("Repeated file name in assets dict: " + - buffer_name); - } else { - throw py::value_error("Asset failed to load: " + buffer_name); - } - } - } - } - - raw::MjModel* model = loadfunc(filename.c_str(), vfs_ptr); - mj_deleteVFS(vfs_ptr); - if (model && !model->buffer) { - mj_deleteModel(model); - model = nullptr; - } - return model; -} - -MjModelWrapper MjModelWrapper::LoadXMLFile( - const std::string& filename, - const std::optional>& assets) { - const auto converted_assets = ConvertAssetsDict(assets); - raw::MjModel* model; - { - py::gil_scoped_release no_gil; - char error[1024]; - model = LoadModelFileImpl( - filename, converted_assets, - [&error](const char* filename, const mjVFS* vfs) { - return InterceptMjErrors(mj_loadXML)( - filename, vfs, error, sizeof(error)); - }); - if (!model) { - throw py::value_error(error); - } - } - return MjModelWrapper(model); -} - -MjModelWrapper MjModelWrapper::LoadBinaryFile( - const std::string& filename, - const std::optional>& assets) { - const auto converted_assets = ConvertAssetsDict(assets); - raw::MjModel* model; - { - py::gil_scoped_release no_gil; - model = LoadModelFileImpl( - filename, converted_assets, InterceptMjErrors(mj_loadModel)); - if (!model) { - throw py::value_error("mj_loadModel: failed to load from mjb"); - } - } - return MjModelWrapper(model); -} - -MjModelWrapper MjModelWrapper::LoadXML( - const std::string& xml, - const std::optional>& assets) { - auto converted_assets = ConvertAssetsDict(assets); - raw::MjModel* model; - { - py::gil_scoped_release no_gil; - std::string model_filename = "model_.xml"; - if (assets.has_value()) { - while (assets->find(model_filename) != assets->end()) { - model_filename = - model_filename.substr(0, model_filename.size() - 4) + "_.xml"; - } - } - converted_assets.emplace_back( - model_filename.c_str(), xml.c_str(), xml.length()); - char error[1024]; - model = LoadModelFileImpl( - model_filename, converted_assets, - [&error](const char* filename, const mjVFS* vfs) { - return InterceptMjErrors(mj_loadXML)( - filename, vfs, error, sizeof(error)); - }); - if (!model) { - throw py::value_error(error); - } - } - return MjModelWrapper(model); -} - -MjModelWrapper MjModelWrapper::WrapRawModel(raw::MjModel* m) { - return MjModelWrapper(m); -} - py::tuple RecompileSpec(raw::MjSpec* spec, const MjModelWrapper& old_m, const MjDataWrapper& old_d) { raw::MjModel* m = static_cast(mju_malloc(sizeof(mjModel))); @@ -428,950 +82,8 @@ py::tuple RecompileSpec(raw::MjSpec* spec, const MjModelWrapper& old_m, return py::make_tuple(m_pyobj, d_pyobj); } -namespace { -// A byte at the start of serialized mjModel structs, which can be incremented -// when we change the serialization logic to reject pickles from an unsupported -// future version. -constexpr static char kSerializationVersion = 1; - -void CheckInput(const std::istream& input, std::string class_name) { - if (input.fail()) { - throw py::value_error("Invalid serialized " + class_name + "."); - } -} - } // namespace -void MjModelWrapper::Serialize(std::ostream& output) const { - WriteChar(output, kSerializationVersion); - - int model_size = mj_sizeModel(get()); - WriteInt(output, model_size); - std::string buffer(model_size, 0); - mj_saveModel(get(), nullptr, buffer.data(), model_size); - WriteBytes(output, buffer.data(), model_size); -} - -std::unique_ptr MjModelWrapper::Deserialize( - std::istream& input) { - CheckInput(input, "mjModel"); - - char serializationVersion = ReadChar(input); - CheckInput(input, "mjModel"); - - if (serializationVersion != kSerializationVersion) { - throw py::value_error("Incompatible serialization version."); - } - - std::size_t model_size = ReadInt(input); - CheckInput(input, "mjModel"); - if (model_size < 0) { - throw py::value_error("Invalid serialized mjModel."); - } - std::string model_bytes(model_size, 0); - ReadBytes(input, model_bytes.data(), model_size); - CheckInput(input, "mjModel"); - - raw::MjModel* model = LoadModelFileImpl( - "model.mjb", - {{"model.mjb", model_bytes.data(), static_cast(model_size)}}, - InterceptMjErrors(mj_loadModel)); - if (!model) { - throw py::value_error("Invalid serialized mjModel."); - } - return std::unique_ptr(new MjModelWrapper(model)); -} - -// ==================== MJCONTACT ============================================== -#define X(var) var(InitPyArray(ptr_->var, owner_)) -MjContactWrapper::MjWrapper() - : WrapperBase(new raw::MjContact{}), - X(pos), - X(frame), - X(friction), - X(solref), - X(solreffriction), - X(solimp), - X(H), - X(geom), - X(flex), - X(elem), - X(vert) {} - -MjContactWrapper::MjWrapper(raw::MjContact* ptr, py::handle owner) - : WrapperBase(ptr, owner), - X(pos), - X(frame), - X(friction), - X(solref), - X(solreffriction), - X(solimp), - X(H), - X(geom), - X(flex), - X(elem), - X(vert) {} -#undef X - -MjContactWrapper::MjWrapper(const MjContactWrapper& other) - : MjContactWrapper() { - *this->ptr_ = *other.ptr_; -} - -MjContactList::MjStructList(raw::MjContact* ptr, int nconmax, - int* ncon, py::handle owner) - : StructListBase(ptr, nconmax, owner, /* lazy = */ true), - ncon_(ncon) {} - -// Slicing -MjContactList::MjStructList(MjContactList& other, py::slice slice) - : StructListBase(other, slice), - ncon_(other.ncon_) {} - -// ==================== MJDATA ================================================= -static void MjDataCapsuleDestructor(PyObject* pyobj) { - mj_deleteData( - static_cast(PyCapsule_GetPointer(pyobj, nullptr))); -} - -absl::flat_hash_map& -MjDataRawPointerMap() { - static auto* hash_map = - new absl::flat_hash_map(); - return *hash_map; -} - -MjDataWrapper* MjDataWrapper::FromRawPointer(raw::MjData* m) noexcept { - try { - auto& map = MjDataRawPointerMap(); - { - py::gil_scoped_acquire gil; - auto found = map.find(m); - return found != map.end() ? found->second : nullptr; - } - } catch (...) { - return nullptr; - } -} - -namespace { -// default timer callback (seconds) -mjtNum GetTime() { - using Clock = std::chrono::steady_clock; - using Seconds = std::chrono::duration; - static const Clock::time_point tm_start = Clock::now(); - return Seconds(Clock::now() - tm_start).count(); -} -} // namespace - -MjDataWrapper::MjWrapper(MjModelWrapper* model) - : WrapperBase(InterceptMjErrors(mj_makeData)(model->get()), - &MjDataCapsuleDestructor), -#undef MJ_M -#define MJ_M(x) model->get()->x -#define X(dtype, var, dim0, dim1) \ - var(InitPyArray(X_ARRAY_SHAPE(model->get()->dim0, dim1), ptr_->var, owner_)), - MJDATA_POINTERS -#undef MJ_M -#define MJ_M(x) (x) -#undef X - - contact(MjContactList(ptr_->contact, NConMax(ptr_), &ptr_->ncon, owner_)), - -#define X(dtype, var, dim0, dim1) var(InitPyArray(ptr_->var, owner_)), - MJDATA_VECTOR -#undef X - model_(model), - model_ref_(py::cast(model_)), - indexer_(ptr_, model_->get(), owner_) { - bool is_newly_inserted = false; - { - py::gil_scoped_acquire gil; - is_newly_inserted = MjDataRawPointerMap().insert({ptr_, this}).second; - } - if (!is_newly_inserted) { - throw UnexpectedError( - "MjDataRawPointerMap already contains this raw mjData*"); - } - - // install default timer if not already installed - { - py::gil_scoped_acquire gil; - if (!mjcb_time) { - mjcb_time = GetTime; - } - } -} - -MjDataWrapper::MjWrapper(const MjDataWrapper& other) - : WrapperBase(other.Copy(), &MjDataCapsuleDestructor), -#undef MJ_M -#define MJ_M(x) other.model_->get()->x -#define X(dtype, var, dim0, dim1) \ - var(InitPyArray(X_ARRAY_SHAPE(other.model_->get()->dim0, dim1), ptr_->var, \ - owner_)), - MJDATA_POINTERS -#undef MJ_M -#define MJ_M(x) (x) -#undef X - - contact(MjContactList(ptr_->contact, NConMax(ptr_), &ptr_->ncon, owner_)), - -#define X(dtype, var, dim0, dim1) var(InitPyArray(ptr_->var, owner_)), - MJDATA_VECTOR -#undef X - model_(other.model_), - model_ref_(other.model_ref_), - indexer_(ptr_, model_->get(), owner_) { - bool is_newly_inserted = false; - { - py::gil_scoped_acquire gil; - is_newly_inserted = MjDataRawPointerMap().insert({ptr_, this}).second; - } - if (!is_newly_inserted) { - throw UnexpectedError( - "MjDataRawPointerMap already contains this raw mjData*"); - } -} - -MjDataWrapper::MjWrapper(MjDataWrapper&& other) - : WrapperBase(other.ptr_, other.owner_), -#undef MJ_M -#define MJ_M(x) other.model_->get()->x -#define X(dtype, var, dim0, dim1) \ - var(InitPyArray(X_ARRAY_SHAPE(other.model_->get()->dim0, dim1), ptr_->var, \ - owner_)), - MJDATA_POINTERS -#undef MJ_M -#define MJ_M(x) (x) -#undef X - - contact(MjContactList(ptr_->contact, NConMax(ptr_), &ptr_->ncon, owner_)), - -#define X(dtype, var, dim0, dim1) var(InitPyArray(ptr_->var, owner_)), - MJDATA_VECTOR -#undef X - model_(other.model_), - model_ref_(std::move(other.model_ref_)), - indexer_(ptr_, model_->get(), owner_) { - bool is_newly_inserted = false; - { - py::gil_scoped_acquire gil; - is_newly_inserted = - MjDataRawPointerMap().insert_or_assign(ptr_, this).second; - } - if (is_newly_inserted) { - throw UnexpectedError( - "MjDataRawPointerMap does not contains the moved-from mjData*"); - } - other.ptr_ = nullptr; -} - -MjDataWrapper::MjWrapper(const MjDataWrapper& other, MjModelWrapper* model) - : WrapperBase(other.Copy(), &MjDataCapsuleDestructor), -#undef MJ_M -#define MJ_M(x) other.model_->get()->x -#define X(dtype, var, dim0, dim1) \ - var(InitPyArray(X_ARRAY_SHAPE(other.model_->get()->dim0, dim1), ptr_->var, \ - owner_)), - MJDATA_POINTERS -#undef MJ_M -#define MJ_M(x) (x) -#undef X - - contact(MjContactList(ptr_->contact, NConMax(ptr_), &ptr_->ncon, owner_)), - -#define X(dtype, var, dim0, dim1) var(InitPyArray(ptr_->var, owner_)), - MJDATA_VECTOR -#undef X - model_(model), - model_ref_(py::cast(model_)), - indexer_(ptr_, model_->get(), owner_) { - bool is_newly_inserted = false; - { - py::gil_scoped_acquire gil; - is_newly_inserted = MjDataRawPointerMap().insert({ptr_, this}).second; - } - if (!is_newly_inserted) { - throw UnexpectedError( - "MjDataRawPointerMap already contains this raw mjData*"); - } -} - -MjDataWrapper::MjWrapper(MjModelWrapper* model, raw::MjData* d) - : WrapperBase(d, &MjDataCapsuleDestructor), -#undef MJ_M -#define MJ_M(x) model->get()->x -#define X(dtype, var, dim0, dim1) \ - var(InitPyArray(X_ARRAY_SHAPE(model->get()->dim0, dim1), ptr_->var, owner_)), - MJDATA_POINTERS -#undef MJ_M -#define MJ_M(x) (x) -#undef X - - contact(MjContactList(ptr_->contact, NConMax(ptr_), &ptr_->ncon, owner_)), - -#define X(dtype, var, dim0, dim1) var(InitPyArray(ptr_->var, owner_)), - MJDATA_VECTOR -#undef X - model_(model), - model_ref_(py::cast(model_)), - indexer_(ptr_, model_->get(), owner_) { - bool is_newly_inserted = false; - { - py::gil_scoped_acquire gil; - is_newly_inserted = MjDataRawPointerMap().insert({ptr_, this}).second; - } - if (!is_newly_inserted) { - throw UnexpectedError( - "MjDataRawPointerMap already contains this raw mjData*"); - } -} - -MjDataWrapper::~MjWrapper() { - if (ptr_) { - bool erased = false; - { - py::gil_scoped_acquire gil; - erased = MjDataRawPointerMap().erase(ptr_); - } - if (!erased) { - std::cerr << "MjDataRawPointerMap does not contain this raw mjData*" - << std::endl; - std::terminate(); - } - } -} - -void MjDataWrapper::Serialize(std::ostream& output) const { - // TODO: Replace this custom serialization with a protobuf - WriteChar(output, kSerializationVersion); - - model_->Serialize(output); - - // Write struct and scalar fields -#define X(var) WriteBytes(output, &ptr_->var, sizeof(ptr_->var)) - X(maxuse_stack); - X(maxuse_arena); - X(maxuse_con); - X(maxuse_efc); - X(solver); - X(timer); - X(warning); - X(ncon); - X(ne); - X(nf); - X(nJ); - X(nA); - X(nefc); - X(nisland); - X(time); - X(energy); -#undef X - - // Write buffer and arena contents - { - MJDATA_POINTERS_PREAMBLE((this->model_->get())) - -#define X(type, name, nr, nc) \ - WriteBytes(output, ptr_->name, sizeof(type)*(this->model_->get()->nr)*(nc)); - MJDATA_POINTERS -#undef X - -#undef MJ_M -#define MJ_M(x) this->model_->get()->x -#undef MJ_D -#define MJ_D(x) this->ptr_->x -#define X(type, name, nr, nc) \ - if ((nr) * (nc)) { \ - WriteBytes(output, ptr_->name, \ - ptr_->name ? sizeof(type) * (nr) * (nc) : 0); \ - } - - MJDATA_ARENA_POINTERS_CONTACT - MJDATA_ARENA_POINTERS_SOLVER - if (mj_isDual(this->model_->get())) { - MJDATA_ARENA_POINTERS_DUAL - } - if (this->ptr_->nisland) { - MJDATA_ARENA_POINTERS_ISLAND - } -#undef MJ_M -#define MJ_M(x) x -#undef MJ_D -#define MJ_D(x) x -#undef X - } -} - -MjDataWrapper MjDataWrapper::Deserialize(std::istream& input) { - char serializationVersion = ReadChar(input); - CheckInput(input, "mjData"); - - if (serializationVersion != kSerializationVersion) { - throw py::value_error("Incompatible serialization version."); - } - - // Read the model that was used to create the mjData. - std::unique_ptr m_wrapper = - MjModelWrapper::Deserialize(input); - raw::MjModel& m = *m_wrapper->get(); - - bool is_dual = mj_isDual(&m); - - raw::MjData* d = mj_makeData(&m); - if (!d) { - throw py::value_error("Failed to create mjData."); - } - - // Read structs and scalar fields -#define X(var) \ - ReadBytes(input, (void*) &d->var, sizeof(d->var)); \ - CheckInput(input, "mjData"); - - X(maxuse_stack); - X(maxuse_arena); - X(maxuse_con); - X(maxuse_efc); - X(solver); - X(timer); - X(warning); - X(ncon); - X(ne); - X(nf); - X(nJ); - X(nA); - X(nefc); - X(nisland); - X(time); - X(energy); -#undef X - - // Read buffer and arena contents - { - MJDATA_POINTERS_PREAMBLE((&m)) - -#define X(type, name, nr, nc) \ - ReadBytes(input, d->name, sizeof(type)*(m.nr)*(nc)); - MJDATA_POINTERS -#undef X - -#undef MJ_M -#define MJ_M(x) m.x -#undef MJ_D -#define MJ_D(x) d->x -// arena pointers might be null, so we need to check the size before allocating. -#define X(type, name, nr, nc) \ - if ((nr) * (nc)) { \ - std::size_t actual_nbytes = ReadInt(input); \ - if (actual_nbytes) { \ - if (actual_nbytes != sizeof(type) * (nr) * (nc)) { \ - input.setstate(input.rdstate() | std::ios_base::failbit); \ - } else { \ - d->name = static_castname)>( \ - mj_arenaAllocByte(d, sizeof(type) * (nr) * (nc), alignof(type))); \ - input.read(reinterpret_cast(d->name), actual_nbytes); \ - } \ - } else { \ - d->name = nullptr; \ - } \ - } - - MJDATA_ARENA_POINTERS_CONTACT - MJDATA_ARENA_POINTERS_SOLVER - if (is_dual) { - MJDATA_ARENA_POINTERS_DUAL - } - if (d->nisland) { - MJDATA_ARENA_POINTERS_ISLAND - } -#undef MJ_M -#define MJ_M(x) x -#undef MJ_D -#define MJ_D(x) x -#undef X - } - CheckInput(input, "mjData"); - - // All bytes should have been used. - input.ignore(1); - if (!input.eof()) { - throw py::value_error("Invalid serialized mjData."); - } - - return MjDataWrapper(m_wrapper.release(), d); -} - -raw::MjData* MjDataWrapper::Copy() const { - const raw::MjModel* m = model_->get(); - return InterceptMjErrors(mj_copyData)(NULL, m, this->ptr_); -} - -// ==================== MJSTATISTIC ============================================ -#define X(var) var(InitPyArray(ptr_->var, owner_)) -MjStatisticWrapper::MjWrapper() - : WrapperBase(new raw::MjStatistic{}), - X(center) {} - -MjStatisticWrapper::MjWrapper(raw::MjStatistic* ptr, py::handle owner) - : WrapperBase(ptr, owner), - X(center) {} -#undef X - -MjStatisticWrapper::MjWrapper(const MjStatisticWrapper& other) - : MjStatisticWrapper() { - *this->ptr_ = *other.ptr_; -} - -// ==================== MJWARNINGSTAT ========================================== -MjWarningStatWrapper::MjWrapper() - : WrapperBase(new raw::MjWarningStat{}) {} - -MjWarningStatWrapper::MjWrapper(raw::MjWarningStat* ptr, py::handle owner) - : WrapperBase(ptr, owner) {} - -MjWarningStatWrapper::MjWrapper(const MjWarningStatWrapper& other) - : MjWarningStatWrapper() { - *this->ptr_ = *other.ptr_; -} - -#define X(type, var) \ - var(std::vector{num}, std::vector{sizeof(raw::MjWarningStat)}, \ - &ptr->var, owner) -MjWarningStatList::MjStructList(raw::MjWarningStat* ptr, int num, - py::handle owner) - : StructListBase(ptr, num, owner), X(int, lastinfo), X(int, number) {} -#undef X - -// Slicing -#define X(type, var) var(other.var[slice]) -MjWarningStatList::MjStructList(MjWarningStatList& other, py::slice slice) - : StructListBase(other, slice), - X(int, lastinfo), - X(int, number) {} -#undef X - -// ==================== MJTIMERSTAT ============================================ -MjTimerStatWrapper::MjWrapper() - : WrapperBase(new raw::MjTimerStat{}) {} - -MjTimerStatWrapper::MjWrapper(raw::MjTimerStat* ptr, py::handle owner) - : WrapperBase(ptr, owner) {} - -MjTimerStatWrapper::MjWrapper(const MjTimerStatWrapper& other) - : MjTimerStatWrapper() { - *this->ptr_ = *other.ptr_; -} - -#define X(type, var) \ - var(std::vector{num}, std::vector{sizeof(raw::MjTimerStat)}, \ - &ptr->var, owner) -MjTimerStatList::MjStructList(raw::MjTimerStat* ptr, int num, py::handle owner) - : StructListBase(ptr, num, owner), - X(mjtNum, duration), - X(int, number) {} -#undef X - -// Slicing -#define X(type, var) var(other.var[slice]) -MjTimerStatList::MjStructList(MjTimerStatList& other, py::slice slice) - : StructListBase(other, slice), - X(mjtNum, duration), - X(int, number) {} -#undef X - -// ==================== MJSOLVERSTAT =========================================== -MjSolverStatWrapper::MjWrapper() - : WrapperBase(new raw::MjSolverStat{}) {} - -MjSolverStatWrapper::MjWrapper(raw::MjSolverStat* ptr, py::handle owner) - : WrapperBase(ptr, owner) {} - -MjSolverStatWrapper::MjWrapper(const MjSolverStatWrapper& other) - : MjSolverStatWrapper() { - *this->ptr_ = *other.ptr_; -} - -#define X(type, var) \ - var(std::vector{num}, std::vector{sizeof(raw::MjSolverStat)}, \ - &ptr->var, owner) -MjSolverStatList::MjStructList(raw::MjSolverStat* ptr, int num, - py::handle owner) - : StructListBase(ptr, num, owner), - X(mjtNum, improvement), - X(mjtNum, gradient), - X(mjtNum, lineslope), - X(int, nactive), - X(int, nchange), - X(int, neval), - X(int, nupdate) {} -#undef X -#undef XN - -// Slicing -#define X(type, var) var(other.var[slice]) -MjSolverStatList::MjStructList(MjSolverStatList& other, py::slice slice) - : StructListBase(other, slice), - X(mjtNum, improvement), - X(mjtNum, gradient), - X(mjtNum, lineslope), - X(int, nactive), - X(int, nchange), - X(int, neval), - X(int, nupdate) {} -#undef X - -// ==================== MJVPERTURB ============================================= -#define X(var) var(InitPyArray(ptr_->var, owner_)) -MjvPerturbWrapper::MjWrapper() - : WrapperBase([]() { - raw::MjvPerturb* const pert = new raw::MjvPerturb; - mjv_defaultPerturb(pert); - return pert; - }()), - X(refpos), - X(refquat), - X(refselpos), - X(localpos) {} -#undef X - -MjvPerturbWrapper::MjWrapper(const MjvPerturbWrapper& other) - : MjvPerturbWrapper() { - *this->ptr_ = *other.ptr_; -} - -// ==================== MJVCAMERA ============================================== -#define X(var) var(InitPyArray(ptr_->var, owner_)) -MjvCameraWrapper::MjWrapper() - : WrapperBase([]() { - raw::MjvCamera* const cam = new raw::MjvCamera; - mjv_defaultCamera(cam); - return cam; - }()), - X(lookat) {} -#undef X - -MjvCameraWrapper::MjWrapper(const MjvCameraWrapper& other) - : MjvCameraWrapper() { - *this->ptr_ = *other.ptr_; -} - -// ==================== MJVGLCAMERA ============================================ -#define X(var) var(InitPyArray(ptr_->var, owner_)) -MjvGLCameraWrapper::MjWrapper() - : WrapperBase(new raw::MjvGLCamera{}), - X(pos), - X(forward), - X(up) {} - -MjvGLCameraWrapper::MjWrapper(raw::MjvGLCamera* ptr, py::handle owner) - : WrapperBase(ptr, owner), - X(pos), - X(forward), - X(up) {} -#undef X - -MjvGLCameraWrapper::MjWrapper(raw::MjvGLCamera&& other) - : MjvGLCameraWrapper() { - *this->ptr_ = other; -} - -MjvGLCameraWrapper::MjWrapper(const MjvGLCameraWrapper& other) - : MjvGLCameraWrapper() { - *this->ptr_ = *other.ptr_; -} - -// ==================== MJVGEOM ================================================ -#define X(var) var(InitPyArray(ptr_->var, owner_)) -MjvGeomWrapper::MjWrapper() - : WrapperBase(new raw::MjvGeom{}), - X(size), - X(pos), - mat([this]() { - static_assert(sizeof(ptr_->mat) == sizeof(ptr_->mat[0])*9); - return InitPyArray(std::array{3, 3}, ptr_->mat, owner_); - }()), - X(rgba) { - mjv_initGeom(ptr_, mjGEOM_NONE, nullptr, nullptr, nullptr, nullptr); -} - -MjvGeomWrapper::MjWrapper(raw::MjvGeom* ptr, py::handle owner) - : WrapperBase(ptr, owner), - X(size), - X(pos), - mat([this]() { - static_assert(sizeof(ptr_->mat) == sizeof(ptr_->mat[0])*9); - return InitPyArray(std::array{3, 3}, ptr_->mat, owner_); - }()), - X(rgba) {} -#undef X - -MjvGeomWrapper::MjWrapper(const MjvGeomWrapper& other) - : MjvGeomWrapper() { - *this->ptr_ = *other.ptr_; -} - -// ==================== MJVLIGHT =============================================== -#define X(var) var(InitPyArray(ptr_->var, owner_)) -MjvLightWrapper::MjWrapper() - : WrapperBase(new raw::MjvLight{}), - X(pos), - X(dir), - X(attenuation), - X(ambient), - X(diffuse), - X(specular) {} - -MjvLightWrapper::MjWrapper(raw::MjvLight* ptr, py::handle owner) - : WrapperBase(ptr, owner), - X(pos), - X(dir), - X(attenuation), - X(ambient), - X(diffuse), - X(specular) {} -#undef X - -MjvLightWrapper::MjWrapper(const MjvLightWrapper& other) - : MjvLightWrapper() { - *this->ptr_ = *other.ptr_; -} - -// ==================== MJVOPTION ============================================== -#define X(var) var(InitPyArray(ptr_->var, owner_)) -MjvOptionWrapper::MjWrapper() - : WrapperBase([]() { - raw::MjvOption* const opt = new raw::MjvOption; - mjv_defaultOption(opt); - return opt; - }()), - X(geomgroup), - X(sitegroup), - X(jointgroup), - X(tendongroup), - X(actuatorgroup), - X(flexgroup), - X(skingroup), - X(flags) {} -#undef X - -MjvOptionWrapper::MjWrapper(const MjvOptionWrapper& other) - : MjvOptionWrapper() { - *this->ptr_ = *other.ptr_; -} - -// ==================== MJVSCENE =============================================== -static void MjvSceneCapsuleDestructor(PyObject* pyobj) { - py::gil_scoped_acquire gil; - auto* scn = static_cast(PyCapsule_GetPointer(pyobj, nullptr)); - if (scn) { - mjv_freeScene(scn); - delete scn; - } -} - -#define X(var) var(InitPyArray(ptr_->var, owner_)) -#define XN(var, n) var(InitPyArray(std::array{n}, ptr_->var, owner_)) -MjvSceneWrapper::MjWrapper() - : WrapperBase([]() { - raw::MjvScene *const scn = new raw::MjvScene; - mjv_defaultScene(scn); - InterceptMjErrors(mjv_makeScene)(nullptr, scn, 0); - return scn; - }(), &MjvSceneCapsuleDestructor), - nskinvert(0), - XN(geoms, 0), - XN(geomorder, 0), - XN(flexedgeadr, 0), - XN(flexedgenum, 0), - XN(flexvertadr, 0), - XN(flexvertnum, 0), - XN(flexfaceadr, 0), - XN(flexfacenum, 0), - XN(flexfaceused, 0), - XN(flexedge, 0), - XN(flexvert, 0), - XN(flexface, 0), - XN(flexnormal, 0), - XN(flextexcoord, 0), - XN(skinfacenum, 0), - XN(skinvertadr, 0), - XN(skinvertnum, 0), - XN(skinvert, 0), - XN(skinnormal, 0), - X(lights), - X(camera), - X(translate), - X(rotate), - X(flags), - X(framergb) {} - -#define XN(var, n) var(InitPyArray(std::array{n}, ptr_->var, owner_)) -MjvSceneWrapper::MjWrapper(const MjModelWrapper& model, int maxgeom) - : WrapperBase( - [maxgeom](const raw::MjModel* m) { - raw::MjvScene *const scn = new raw::MjvScene; - mjv_defaultScene(scn); - InterceptMjErrors(mjv_makeScene)(m, scn, maxgeom); - return scn; - }(model.get()), - &MjvSceneCapsuleDestructor), - nskinvert([](const raw::MjModel* m) { - int nskinvert = 0; - for (int i = 0; i < m->nskin; ++i) { - nskinvert += m->skin_vertnum[i]; - } - return nskinvert; - }(model.get())), - nflexface([](const raw::MjModel* m) { - int nflexface = 0; - int flexfacenum = 0; - for (int f=0; f < m->nflex; f++) { - if (m->flex_dim[f] == 0) { - // 1D : 0 - flexfacenum = 0; - } else if (m->flex_dim[f] == 2) { - // 2D: 2*fragments + 2*elements - flexfacenum = 2*m->flex_shellnum[f] + 2*m->flex_elemnum[f]; - } else { - // 3D: max(fragments, 4*maxlayer) - // find number of elements in biggest layer - int maxlayer = 0, layer = 0, nlayer = 1; - while (nlayer) { - nlayer = 0; - for (int e=0; e < m->flex_elemnum[f]; e++) { - if (m->flex_elemlayer[m->flex_elemadr[f]+e] == layer) { - nlayer++; - } - } - maxlayer = mjMAX(maxlayer, nlayer); - layer++; - } - flexfacenum = mjMAX(m->flex_shellnum[f], 4*maxlayer); - } - - // accumulate over flexes - nflexface += flexfacenum; - } - return nflexface; - }(model.get())), - nflexedge(model.get()->nflexedge), - nflexvert(model.get()->nflexvert), - XN(geoms, ptr_->maxgeom), - XN(geomorder, ptr_->maxgeom), - XN(flexedgeadr, ptr_->nflex), - XN(flexedgenum, ptr_->nflex), - XN(flexvertadr, ptr_->nflex), - XN(flexvertnum, ptr_->nflex), - XN(flexfaceadr, ptr_->nflex), - XN(flexfacenum, ptr_->nflex), - XN(flexfaceused, ptr_->nflex), - XN(flexedge, 2*nflexedge), - XN(flexvert, 3*nflexvert), - XN(flexface, 9*nflexface), - XN(flexnormal, 9*nflexface), - XN(flextexcoord, 6*nflexface), - XN(skinfacenum, ptr_->nskin), - XN(skinvertadr, ptr_->nskin), - XN(skinvertnum, ptr_->nskin), - XN(skinvert, 3*nskinvert), - XN(skinnormal, 3*nskinvert), - X(lights), - X(camera), - X(translate), - X(rotate), - X(flags), - X(framergb) {} -#undef X -#undef XN - -template -static T* MallocAndCopy(const T* src, int count) { - if (src) { - T* out = static_cast(mju_malloc(count * sizeof(T))); - std::memcpy(out, src, count * sizeof(T)); - return out; - } else { - return nullptr; - } -} - -MjvSceneWrapper::MjWrapper(const MjvSceneWrapper& other) - : MjvSceneWrapper() { - mjv_freeScene(ptr_); - *ptr_ = *other.ptr_; - -#define XN(var, n) \ - ptr_->var = MallocAndCopy(other.ptr_->var, n); \ - var = InitPyArray(std::array{n}, ptr_->var, owner_); - - XN(geoms, ptr_->ngeom); - XN(geomorder, ptr_->ngeom); - XN(flexedgeadr, ptr_->nflex); - XN(flexedgenum, ptr_->nflex); - XN(flexvertadr, ptr_->nflex); - XN(flexvertnum, ptr_->nflex); - XN(flexfaceadr, ptr_->nflex); - XN(flexfacenum, ptr_->nflex); - XN(flexfaceused, ptr_->nflex); - XN(flexedge, 2*nflexedge); - XN(flexvert, 3*nflexvert); - XN(flexface, 9*nflexface); - XN(flexnormal, 9*nflexface); - XN(flextexcoord, 6*nflexface); - XN(skinfacenum, ptr_->nskin); - XN(skinvertadr, ptr_->nskin); - XN(skinvertnum, ptr_->nskin); - XN(skinvert, 3*nskinvert); - XN(skinnormal, 3*nskinvert); - -#undef XN -} - -// ==================== MJVFIGURE ============================================== -#define X(var) var(InitPyArray(ptr_->var, owner_)) -MjvFigureWrapper::MjWrapper() - : WrapperBase([]() { - raw::MjvFigure* const fig = new raw::MjvFigure; - mjv_defaultFigure(fig); - return fig; - }()), - X(flg_ticklabel), - X(gridsize), - X(gridrgb), - X(figurergba), - X(panergba), - X(legendrgba), - X(textrgb), - X(linergb), - X(range), - X(highlight), - X(linepnt), - X(linedata), - X(xaxispixel), - X(yaxispixel), - X(xaxisdata), - X(yaxisdata), -#undef X - - linename([](raw::MjvFigure* ptr, py::handle owner) { -// Use a macro to help us static_assert that the array extents here are kept -// in sync with mjVisualize.h. -#define MAKE_STR_ARRAY(N1, N2) \ - static_assert( \ - std::is_same_v); \ - return py::array(py::dtype("|S" #N2), N1, ptr->linename, owner); - - MAKE_STR_ARRAY(mjMAXLINE, 100); - -#undef MAKE_STR_ARRAY - }(ptr_, owner_)) {} - -MjvFigureWrapper::MjWrapper(const MjvFigureWrapper& other) - : MjvFigureWrapper() { - *this->ptr_ = *other.ptr_; -} - PYBIND11_MODULE(_structs, m) { py::module_::import("mujoco._enums"); @@ -1439,8 +151,8 @@ PYBIND11_MODULE(_structs, m) { << self.attr("__class__").attr("__name__").cast(); #define X(type, var) \ - result << "\n " #var ": "; \ - StructReprImpl(self.attr(#var), result, 2); + result << "\n " #var ": "; \ + StructReprImpl(self.attr(#var), result, 2); X(raw::MjVisualGlobal, global_) X(raw::MjVisualQuality, quality) @@ -1457,10 +169,10 @@ PYBIND11_MODULE(_structs, m) { mjVisualGlobal.def("__copy__", [](const raw::MjVisualGlobal& other) { return raw::MjVisualGlobal(other); }); - mjVisualGlobal.def( - "__deepcopy__", [](const raw::MjVisualGlobal& other, py::dict) { - return raw::MjVisualGlobal(other); - }); + mjVisualGlobal.def("__deepcopy__", + [](const raw::MjVisualGlobal& other, py::dict) { + return raw::MjVisualGlobal(other); + }); DefineStructFunctions(mjVisualGlobal); #define X(var) mjVisualGlobal.def_readwrite(#var, &raw::MjVisualGlobal::var) X(orthographic); @@ -1481,10 +193,10 @@ PYBIND11_MODULE(_structs, m) { mjVisualQuality.def("__copy__", [](const raw::MjVisualQuality& other) { return raw::MjVisualQuality(other); }); - mjVisualQuality.def( - "__deepcopy__", [](const raw::MjVisualQuality& other, py::dict) { - return raw::MjVisualQuality(other); - }); + mjVisualQuality.def("__deepcopy__", + [](const raw::MjVisualQuality& other, py::dict) { + return raw::MjVisualQuality(other); + }); DefineStructFunctions(mjVisualQuality); #define X(var) mjVisualQuality.def_readwrite(#var, &raw::MjVisualQuality::var) X(shadowsize); @@ -1495,21 +207,20 @@ PYBIND11_MODULE(_structs, m) { #undef X py::class_ mjVisualHeadlight(mjVisual, "Headlight"); - mjVisualHeadlight.def( - "__copy__", [](const MjVisualHeadlightWrapper& other) { - return MjVisualHeadlightWrapper(other); - }); - mjVisualHeadlight.def( - "__deepcopy__", [](const MjVisualHeadlightWrapper& other, py::dict) { - return MjVisualHeadlightWrapper(other); - }); + mjVisualHeadlight.def("__copy__", [](const MjVisualHeadlightWrapper& other) { + return MjVisualHeadlightWrapper(other); + }); + mjVisualHeadlight.def("__deepcopy__", + [](const MjVisualHeadlightWrapper& other, py::dict) { + return MjVisualHeadlightWrapper(other); + }); DefineStructFunctions(mjVisualHeadlight); - #define X(var) \ - DefinePyArray(mjVisualHeadlight, #var, &MjVisualHeadlightWrapper::var) +#define X(var) \ + DefinePyArray(mjVisualHeadlight, #var, &MjVisualHeadlightWrapper::var) X(ambient); X(diffuse); X(specular); - #undef X +#undef X mjVisualHeadlight.def_property( "active", [](const MjVisualHeadlightWrapper& c) { return c.get()->active; }, @@ -1521,10 +232,9 @@ PYBIND11_MODULE(_structs, m) { mjVisualMap.def("__copy__", [](const raw::MjVisualMap& other) { return raw::MjVisualMap(other); }); - mjVisualMap.def( - "__deepcopy__", [](const raw::MjVisualMap& other, py::dict) { - return raw::MjVisualMap(other); - }); + mjVisualMap.def("__deepcopy__", [](const raw::MjVisualMap& other, py::dict) { + return raw::MjVisualMap(other); + }); DefineStructFunctions(mjVisualMap); #define X(var) mjVisualMap.def_readwrite(#var, &raw::MjVisualMap::var) X(stiffness); @@ -1546,10 +256,10 @@ PYBIND11_MODULE(_structs, m) { mjVisualScale.def("__copy__", [](const raw::MjVisualScale& other) { return raw::MjVisualScale(other); }); - mjVisualScale.def( - "__deepcopy__", [](const raw::MjVisualScale& other, py::dict) { - return raw::MjVisualScale(other); - }); + mjVisualScale.def("__deepcopy__", + [](const raw::MjVisualScale& other, py::dict) { + return raw::MjVisualScale(other); + }); DefineStructFunctions(mjVisualScale); #define X(var) mjVisualScale.def_readwrite(#var, &raw::MjVisualScale::var) X(forcewidth); @@ -1575,10 +285,10 @@ PYBIND11_MODULE(_structs, m) { mjVisualRgba.def("__copy__", [](const MjVisualRgbaWrapper& other) { return MjVisualRgbaWrapper(other); }); - mjVisualRgba.def( - "__deepcopy__", [](const MjVisualRgbaWrapper& other, py::dict) { - return MjVisualRgbaWrapper(other); - }); + mjVisualRgba.def("__deepcopy__", + [](const MjVisualRgbaWrapper& other, py::dict) { + return MjVisualRgbaWrapper(other); + }); DefineStructFunctions(mjVisualRgba); #define X(var) DefinePyArray(mjVisualRgba, #var, &MjVisualRgbaWrapper::var) X(fog); @@ -1626,26 +336,26 @@ PYBIND11_MODULE(_structs, m) { // ==================== MJMODEL ============================================== py::class_ mjModel(m, "MjModel"); mjModel.def_static( - "from_xml_string", &MjModelWrapper::LoadXML, - py::arg("xml"), py::arg_v("assets", py::none()), + "from_xml_string", &MjModelWrapper::LoadXML, py::arg("xml"), + py::arg_v("assets", py::none()), py::doc( -R"(Loads an MjModel from an XML string and an optional assets dictionary.)")); + R"(Loads an MjModel from an XML string and an optional assets dictionary.)")); mjModel.def_static("_from_model_ptr", [](uintptr_t addr) { return MjModelWrapper::WrapRawModel(reinterpret_cast(addr)); }); mjModel.def_static( - "from_xml_path", &MjModelWrapper::LoadXMLFile, - py::arg("filename"), py::arg_v("assets", py::none()), + "from_xml_path", &MjModelWrapper::LoadXMLFile, py::arg("filename"), + py::arg_v("assets", py::none()), py::doc( -R"(Loads an MjModel from an XML file and an optional assets dictionary. + R"(Loads an MjModel from an XML file and an optional assets dictionary. The filename for the XML can also refer to a key in the assets dictionary. This is useful for example when the XML is not available as a file on disk.)")); mjModel.def_static( - "from_binary_path", &MjModelWrapper::LoadBinaryFile, - py::arg("filename"), py::arg_v("assets", py::none()), + "from_binary_path", &MjModelWrapper::LoadBinaryFile, py::arg("filename"), + py::arg_v("assets", py::none()), py::doc( -R"(Loads an MjModel from an MJB file and an optional assets dictionary. + R"(Loads an MjModel from an MJB file and an optional assets dictionary. The filename for the MJB can also refer to a key in the assets dictionary. This is useful for example when the MJB is not available as a file on disk.)")); @@ -1705,34 +415,37 @@ This is useful for example when the MJB is not available as a file on disk.)")); return py::tuple(py::cast(fields)); }); -#define X(dtype, var, dim0, dim1) \ - if constexpr (std::string_view(#var) != "text_data" && \ - std::string_view(#var) != "names" && \ - std::string_view(#var) != "paths") { \ - DefinePyArray(mjModel, #var, &MjModelWrapper::var); \ +#define X(dtype, var, dim0, dim1) \ + if constexpr (std::string_view(#var) != "text_data" && \ + std::string_view(#var) != "names" && \ + std::string_view(#var) != "paths") { \ + DefinePyArray(mjModel, #var, &MjModelWrapper::var); \ } MJMODEL_POINTERS #undef X - mjModel.def_property_readonly( - "text_data", [](const MjModelWrapper& m) -> const auto& { - // Return the full bytes array of concatenated text data - return m.text_data_bytes; - }); - mjModel.def_property_readonly( - "names", [](const MjModelWrapper& m) -> const auto& { - // Return the full bytes array of concatenated names - return m.names_bytes; - }); - mjModel.def_property_readonly( - "paths", [](const MjModelWrapper& m) -> const auto& { - // Return the full bytes array of concatenated paths - return m.paths_bytes; - }); - mjModel.def_property_readonly( - "signature", [](const MjModelWrapper& m) -> const uint64_t& { - return m.get()->signature; - }); + mjModel.def_property_readonly("text_data", + [](const MjModelWrapper& m) -> const auto& { + // Return the full bytes array of concatenated + // text data + return m.text_data_bytes; + }); + mjModel.def_property_readonly("names", + [](const MjModelWrapper& m) -> const auto& { + // Return the full bytes array of concatenated + // names + return m.names_bytes; + }); + mjModel.def_property_readonly("paths", + [](const MjModelWrapper& m) -> const auto& { + // Return the full bytes array of concatenated + // paths + return m.paths_bytes; + }); + mjModel.def_property_readonly("signature", + [](const MjModelWrapper& m) -> const uint64_t& { + return m.get()->signature; + }); #define XGROUP(MjModelGroupedViews, field, nfield, FIELD_XMACROS) \ mjModel.def( \ @@ -1740,28 +453,28 @@ This is useful for example when the MJB is not available as a file on disk.)")); [](MjModelWrapper& m, int i) -> auto& { return m.indexer().field(i); }, \ py::return_value_policy::reference_internal); \ mjModel.def( \ - #field, [](MjModelWrapper& m, std::string_view name) -> auto& { \ + #field, \ + [](MjModelWrapper& m, std::string_view name) -> auto& { \ return m.indexer().field##_by_name(name); \ }, \ py::return_value_policy::reference_internal, py::arg_v("name", "")); - MJMODEL_VIEW_GROUPS #undef XGROUP -#define XGROUP(spectype, field) \ - mjModel.def( \ - "bind_scalar", \ - [](MjModelWrapper& m, spectype& spec) -> auto& { \ - if (mjs_getSpec(spec.element)->element->signature != \ - m.get()->signature) { \ - throw py::value_error( \ - "The mjSpec does not match mjModel. Please recompile " \ - "the mjSpec."); \ - } \ - return m.indexer().field(mjs_getId(spec.element)); \ - }, \ - py::return_value_policy::reference_internal, \ +#define XGROUP(spectype, field) \ + mjModel.def( \ + "bind_scalar", \ + [](MjModelWrapper& m, spectype& spec) -> auto& { \ + if (mjs_getSpec(spec.element)->element->signature != \ + m.get()->signature) { \ + throw py::value_error( \ + "The mjSpec does not match mjModel. Please recompile " \ + "the mjSpec."); \ + } \ + return m.indexer().field(mjs_getId(spec.element)); \ + }, \ + py::return_value_policy::reference_internal, \ py::arg_v("spec", py::none())); MJMODEL_BIND_GROUPS @@ -1773,7 +486,8 @@ This is useful for example when the MJB is not available as a file on disk.)")); [](MjModelWrapper& m, int i) -> auto& { return m.indexer().field(i); }, \ py::return_value_policy::reference_internal); \ mjModel.def( \ - #altname, [](MjModelWrapper& m, std::string_view name) -> auto& { \ + #altname, \ + [](MjModelWrapper& m, std::string_view name) -> auto& { \ return m.indexer().field##_by_name(name); \ }, \ py::return_value_policy::reference_internal, py::arg_v("name", "")); @@ -1787,12 +501,10 @@ This is useful for example when the MJB is not available as a file on disk.)")); py::class_ groupedViews(m, "_" #MjModelGroupedViews); \ FIELD_XMACROS \ groupedViews.def("__repr__", MjModelStructRepr); \ - groupedViews.def_property_readonly("id", [](GroupedViews& views) { \ - return views.index(); \ - }); \ - groupedViews.def_property_readonly("name", [](GroupedViews& views) { \ - return views.name(); \ - }); \ + groupedViews.def_property_readonly( \ + "id", [](GroupedViews& views) { return views.index(); }); \ + groupedViews.def_property_readonly( \ + "name", [](GroupedViews& views) { return views.name(); }); \ } #define X(type, prefix, var, dim0, dim1) \ groupedViews.def_property( \ @@ -1807,8 +519,8 @@ This is useful for example when the MJB is not available as a file on disk.)")); { py::handle builtins(PyEval_GetBuiltins()); builtins[MjModelWrapper::kFromRawPointer] = - reinterpret_cast(reinterpret_cast( - &MjModelWrapper::FromRawPointer)); + reinterpret_cast( + reinterpret_cast(&MjModelWrapper::FromRawPointer)); } // ==================== MJWARNINGSTAT ======================================== @@ -1817,10 +529,10 @@ This is useful for example when the MJB is not available as a file on disk.)")); mjWarningStat.def("__copy__", [](const MjWarningStatWrapper& other) { return MjWarningStatWrapper(other); }); - mjWarningStat.def( - "__deepcopy__", [](const MjWarningStatWrapper& other, py::dict) { - return MjWarningStatWrapper(other); - }); + mjWarningStat.def("__deepcopy__", + [](const MjWarningStatWrapper& other, py::dict) { + return MjWarningStatWrapper(other); + }); DefineStructFunctions(mjWarningStat); #define X(var) \ mjWarningStat.def_property( \ @@ -1833,10 +545,8 @@ This is useful for example when the MJB is not available as a file on disk.)")); #undef X py::class_ mjWarningStatList(m, "_MjWarningStatList"); - mjWarningStatList.def( - "__getitem__", - &MjWarningStatList::operator[], - py::return_value_policy::reference); + mjWarningStatList.def("__getitem__", &MjWarningStatList::operator[], + py::return_value_policy::reference); mjWarningStatList.def( "__getitem__", [](MjWarningStatList& list, ::mjtWarning idx) { return list[idx]; }, @@ -1857,10 +567,10 @@ This is useful for example when the MJB is not available as a file on disk.)")); mjTimerStat.def("__copy__", [](const MjTimerStatWrapper& other) { return MjTimerStatWrapper(other); }); - mjTimerStat.def( - "__deepcopy__", [](const MjTimerStatWrapper& other, py::dict) { - return MjTimerStatWrapper(other); - }); + mjTimerStat.def("__deepcopy__", + [](const MjTimerStatWrapper& other, py::dict) { + return MjTimerStatWrapper(other); + }); DefineStructFunctions(mjTimerStat); #define X(var) \ mjTimerStat.def_property( \ @@ -1873,10 +583,8 @@ This is useful for example when the MJB is not available as a file on disk.)")); #undef X py::class_ mjTimerStatList(m, "_MjTimerStatList"); - mjTimerStatList.def( - "__getitem__", - &MjTimerStatList::operator[], - py::return_value_policy::reference); + mjTimerStatList.def("__getitem__", &MjTimerStatList::operator[], + py::return_value_policy::reference); mjTimerStatList.def( "__getitem__", [](MjTimerStatList& list, ::mjtTimer idx) { return list[idx]; }, @@ -1885,8 +593,7 @@ This is useful for example when the MJB is not available as a file on disk.)")); mjTimerStatList.def("__len__", &MjTimerStatList::size); DefineStructFunctions(mjTimerStatList); -#define X(type, var) \ - mjTimerStatList.def_readonly(#var, &MjTimerStatList::var) +#define X(type, var) mjTimerStatList.def_readonly(#var, &MjTimerStatList::var) X(mjtNum, duration); X(int, number); #undef X @@ -1897,10 +604,10 @@ This is useful for example when the MJB is not available as a file on disk.)")); mjSolverStat.def("__copy__", [](const MjSolverStatWrapper& other) { return MjSolverStatWrapper(other); }); - mjSolverStat.def( - "__deepcopy__", [](const MjSolverStatWrapper& other, py::dict) { - return MjSolverStatWrapper(other); - }); + mjSolverStat.def("__deepcopy__", + [](const MjSolverStatWrapper& other, py::dict) { + return MjSolverStatWrapper(other); + }); DefineStructFunctions(mjSolverStat); #define X(var) \ mjSolverStat.def_property( \ @@ -1919,13 +626,12 @@ This is useful for example when the MJB is not available as a file on disk.)")); py::class_ mjSolverStatList(m, "_MjSolverStatList"); mjSolverStatList.def("__getitem__", &MjSolverStatList::operator[], - py::return_value_policy::reference); + py::return_value_policy::reference); mjSolverStatList.def("__getitem__", &MjSolverStatList::Slice); mjSolverStatList.def("__len__", &MjSolverStatList::size); DefineStructFunctions(mjSolverStatList); -#define X(type, var) \ - mjSolverStatList.def_readonly(#var, &MjSolverStatList::var) +#define X(type, var) mjSolverStatList.def_readonly(#var, &MjSolverStatList::var) X(mjtNum, improvement); X(mjtNum, gradient); X(mjtNum, lineslope); @@ -2025,12 +731,10 @@ This is useful for example when the MJB is not available as a file on disk.)")); mjData.def_property_readonly("_address", [](const MjDataWrapper& d) { return reinterpret_cast(d.get()); }); - mjData.def_property_readonly("model", [](const MjDataWrapper& d) { - return &d.model(); - }); - mjData.def("__copy__", [](const MjDataWrapper& other) { - return MjDataWrapper(other); - }); + mjData.def_property_readonly( + "model", [](const MjDataWrapper& d) { return &d.model(); }); + mjData.def("__copy__", + [](const MjDataWrapper& other) { return MjDataWrapper(other); }); mjData.def("__deepcopy__", [](const MjDataWrapper& other, py::dict memo) { // Use copy.deepcopy(model) to make a model that Python is aware of. py::object new_model_py = @@ -2048,9 +752,8 @@ This is useful for example when the MJB is not available as a file on disk.)")); return MjDataWrapper::Deserialize(input); })); mjData.def_property_readonly( - "signature", [](const MjDataWrapper& d) -> uint64_t { - return d.get()->signature; - }); + "signature", + [](const MjDataWrapper& d) -> uint64_t { return d.get()->signature; }); #define X(type, var) \ mjData.def_property( \ @@ -2097,7 +800,8 @@ This is useful for example when the MJB is not available as a file on disk.)")); [](MjDataWrapper& d, int i) -> auto& { return d.indexer().field(i); }, \ py::return_value_policy::reference_internal); \ mjData.def( \ - #field, [](MjDataWrapper& d, std::string_view name) -> auto& { \ + #field, \ + [](MjDataWrapper& d, std::string_view name) -> auto& { \ return d.indexer().field##_by_name(name); \ }, \ py::return_value_policy::reference_internal, py::arg_v("name", "")); @@ -2105,19 +809,19 @@ This is useful for example when the MJB is not available as a file on disk.)")); MJDATA_VIEW_GROUPS #undef XGROUP -#define XGROUP(spectype, field) \ - mjData.def( \ - "bind_scalar", \ - [](MjDataWrapper& d, spectype& spec) -> auto& { \ - if (mjs_getSpec(spec.element)->element->signature != \ - d.get()->signature) { \ - throw py::value_error( \ - "The mjSpec does not match mjData. Please recompile "\ - "the mjSpec."); \ - } \ - return d.indexer().field(mjs_getId(spec.element)); \ - }, \ - py::return_value_policy::reference_internal, \ +#define XGROUP(spectype, field) \ + mjData.def( \ + "bind_scalar", \ + [](MjDataWrapper& d, spectype& spec) -> auto& { \ + if (mjs_getSpec(spec.element)->element->signature != \ + d.get()->signature) { \ + throw py::value_error( \ + "The mjSpec does not match mjData. Please recompile " \ + "the mjSpec."); \ + } \ + return d.indexer().field(mjs_getId(spec.element)); \ + }, \ + py::return_value_policy::reference_internal, \ py::arg_v("spec", py::none())); MJDATA_BIND_GROUPS @@ -2129,7 +833,8 @@ This is useful for example when the MJB is not available as a file on disk.)")); [](MjDataWrapper& d, int i) -> auto& { return d.indexer().field(i); }, \ py::return_value_policy::reference_internal); \ mjData.def( \ - #altname, [](MjDataWrapper& d, std::string_view name) -> auto& { \ + #altname, \ + [](MjDataWrapper& d, std::string_view name) -> auto& { \ return d.indexer().field##_by_name(name); \ }, \ py::return_value_policy::reference_internal, py::arg_v("name", "")); @@ -2143,12 +848,10 @@ This is useful for example when the MJB is not available as a file on disk.)")); py::class_ groupedViews(m, "_" #MjDataGroupedViews); \ FIELD_XMACROS \ groupedViews.def("__repr__", MjDataStructRepr); \ - groupedViews.def_property_readonly("id", [](GroupedViews& views) { \ - return views.index(); \ - }); \ - groupedViews.def_property_readonly("name", [](GroupedViews& views) { \ - return views.name(); \ - }); \ + groupedViews.def_property_readonly( \ + "id", [](GroupedViews& views) { return views.index(); }); \ + groupedViews.def_property_readonly( \ + "name", [](GroupedViews& views) { return views.name(); }); \ } #define X(type, prefix, var, dim0, dim1) \ groupedViews.def_property( \ @@ -2162,9 +865,8 @@ This is useful for example when the MJB is not available as a file on disk.)")); { py::handle builtins(PyEval_GetBuiltins()); - builtins[MjDataWrapper::kFromRawPointer] = - reinterpret_cast(reinterpret_cast( - &MjDataWrapper::FromRawPointer)); + builtins[MjDataWrapper::kFromRawPointer] = reinterpret_cast( + reinterpret_cast(&MjDataWrapper::FromRawPointer)); } // ==================== MJSTATISTIC ========================================== @@ -2173,10 +875,10 @@ This is useful for example when the MJB is not available as a file on disk.)")); mjStatistic.def("__copy__", [](const MjStatisticWrapper& other) { return MjStatisticWrapper(other); }); - mjStatistic.def( - "__deepcopy__", [](const MjStatisticWrapper& other, py::dict) { - return MjStatisticWrapper(other); - }); + mjStatistic.def("__deepcopy__", + [](const MjStatisticWrapper& other, py::dict) { + return MjStatisticWrapper(other); + }); DefineStructFunctions(mjStatistic); #define X(var) \ @@ -2198,9 +900,8 @@ This is useful for example when the MJB is not available as a file on disk.)")); // ==================== MJLROPT ============================================== py::class_ mjLROpt(m, "MjLROpt"); mjLROpt.def(py::init<>()); - mjLROpt.def("__copy__", [](const raw::MjLROpt& other) { - return raw::MjLROpt(other); - }); + mjLROpt.def("__copy__", + [](const raw::MjLROpt& other) { return raw::MjLROpt(other); }); mjLROpt.def("__deepcopy__", [](const raw::MjLROpt& other, py::dict) { return raw::MjLROpt(other); }); @@ -2285,11 +986,10 @@ This is useful for example when the MJB is not available as a file on disk.)")); mjvGLCamera.def("__copy__", [](const MjvGLCameraWrapper& other) { return MjvGLCameraWrapper(other); }); - mjvGLCamera.def( - "__deepcopy__", - [](const MjvGLCameraWrapper& other, py::dict) { - return MjvGLCameraWrapper(other); - }); + mjvGLCamera.def("__deepcopy__", + [](const MjvGLCameraWrapper& other, py::dict) { + return MjvGLCameraWrapper(other); + }); DefineStructFunctions(mjvGLCamera); #define X(var) \ mjvGLCamera.def_property( \ @@ -2400,7 +1100,7 @@ This is useful for example when the MJB is not available as a file on disk.)")); #define X(var) \ mjvOption.def_property( \ #var, [](const MjvOptionWrapper& c) { return c.get()->var; }, \ - [](MjvOptionWrapper& c, decltype(raw::MjvOption::var) rhs) { \ + [](MjvOptionWrapper& c, decltype(raw::MjvOption::var) rhs) { \ c.get()->var = rhs; \ }) X(label); @@ -2423,8 +1123,8 @@ This is useful for example when the MJB is not available as a file on disk.)")); // ==================== MJVSCENE ============================================= py::class_ mjvScene(m, "MjvScene"); mjvScene.def(py::init<>()); - mjvScene.def(py::init(), - py::arg("model"), py::arg("maxgeom")); + mjvScene.def(py::init(), py::arg("model"), + py::arg("maxgeom")); mjvScene.def("__copy__", [](const MjvSceneWrapper& other) { return MjvSceneWrapper(other); }); @@ -2551,11 +1251,9 @@ This is useful for example when the MJB is not available as a file on disk.)")); py::arg("cam1"), py::arg("cam2"), py::doc(python_traits::mjv_averageCamera::doc)); - m.def( - "_recompile_spec_addr", - [](uintptr_t spec_addr, const MjModelWrapper& m, const MjDataWrapper& d) { - return RecompileSpec(reinterpret_cast(spec_addr), m, d); - } - ); + m.def("_recompile_spec_addr", [](uintptr_t spec_addr, const MjModelWrapper& m, + const MjDataWrapper& d) { + return RecompileSpec(reinterpret_cast(spec_addr), m, d); + }); } // PYBIND11_MODULE NOLINT(readability/fn_size) } // namespace mujoco::python::_impl diff --git a/python/mujoco/structs_wrappers.cc b/python/mujoco/structs_wrappers.cc new file mode 100644 index 00000000..7c1ec9b0 --- /dev/null +++ b/python/mujoco/structs_wrappers.cc @@ -0,0 +1,1328 @@ +// 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 +#include +#include +#include // NOLINT(build/c++11) +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include + +#include +#include +#include +#include "errors.h" +#include "private.h" +#include "raw.h" +#include "serialization.h" +#include "structs.h" +#include +#include +#include +#include +#include +#include +#include + +namespace mujoco::python::_impl { + +namespace py = ::pybind11; + +namespace { +#define PTRDIFF(x, y) \ + reinterpret_cast(x) - reinterpret_cast(y) + +// Returns the shape of a NumPy array given the dimensions from an X Macro. +// If dim1 is a _literal_ constant 1, the resulting array is 1-dimensional of +// length dim0, otherwise the resulting array is 2-dimensional of shape +// (dim0, dim1). +#define X_ARRAY_SHAPE(dim0, dim1) XArrayShapeImpl(#dim1)((dim0), (dim1)) + +std::vector XArrayShapeImpl1D(int dim0, int dim1) { return {dim0}; } + +std::vector XArrayShapeImpl2D(int dim0, int dim1) { return {dim0, dim1}; } + +constexpr auto XArrayShapeImpl(const std::string_view dim1_str) { + if (dim1_str == "1") { + return XArrayShapeImpl1D; + } else { + return XArrayShapeImpl2D; + } +} + +inline std::size_t NConMax(const mjData* d) { + return d->narena / sizeof(mjContact); +} + +} // namespace + +// ==================== MJOPTION =============================================== +#define X(var, dim) , var(InitPyArray(std::array{dim}, ptr_->var, owner_)) +MjOptionWrapper::MjWrapper() + : WrapperBase([]() { + raw::MjOption* const opt = new raw::MjOption; + mj_defaultOption(opt); + return opt; + }()) MJOPTION_VECTORS {} + +MjOptionWrapper::MjWrapper(raw::MjOption* ptr, py::handle owner) + : WrapperBase(ptr, owner) MJOPTION_VECTORS {} +#undef X + +MjOptionWrapper::MjWrapper(const MjOptionWrapper& other) : MjOptionWrapper() { + *this->ptr_ = *other.ptr_; +} + +// ==================== MJVISUAL =============================================== +#define X(var) var(InitPyArray(ptr_->var, owner_)) +MjVisualHeadlightWrapper::MjWrapper() + : WrapperBase(new raw::MjVisualHeadlight{}), + X(ambient), + X(diffuse), + X(specular) {} + +MjVisualHeadlightWrapper::MjWrapper(raw::MjVisualHeadlight* ptr, + py::handle owner) + : WrapperBase(ptr, owner), X(ambient), X(diffuse), X(specular) {} +#undef X + +MjVisualHeadlightWrapper::MjWrapper(const MjVisualHeadlightWrapper& other) + : MjVisualHeadlightWrapper() { + *this->ptr_ = *other.ptr_; +} + +#define X(var) var(InitPyArray(ptr_->var, owner_)) +MjVisualRgbaWrapper::MjWrapper() + : WrapperBase(new raw::MjVisualRgba{}), + X(fog), + X(haze), + X(force), + X(inertia), + X(joint), + X(actuator), + X(actuatornegative), + X(actuatorpositive), + X(com), + X(camera), + X(light), + X(selectpoint), + X(connect), + X(contactpoint), + X(contactforce), + X(contactfriction), + X(contacttorque), + X(contactgap), + X(rangefinder), + X(constraint), + X(slidercrank), + X(crankbroken), + X(frustum) {} + +MjVisualRgbaWrapper::MjWrapper(raw::MjVisualRgba* ptr, py::handle owner) + : WrapperBase(ptr, owner), + X(fog), + X(haze), + X(force), + X(inertia), + X(joint), + X(actuator), + X(actuatornegative), + X(actuatorpositive), + X(com), + X(camera), + X(light), + X(selectpoint), + X(connect), + X(contactpoint), + X(contactforce), + X(contactfriction), + X(contacttorque), + X(contactgap), + X(rangefinder), + X(constraint), + X(slidercrank), + X(crankbroken), + X(frustum) {} +#undef X + +MjVisualRgbaWrapper::MjWrapper(const MjVisualRgbaWrapper& other) + : MjVisualRgbaWrapper() { + *this->ptr_ = *other.ptr_; +} + +MjVisualWrapper::MjWrapper() + : WrapperBase(new raw::MjVisual{}), + headlight(&ptr_->headlight, owner_), + rgba(&ptr_->rgba, owner_) {} + +MjVisualWrapper::MjWrapper(raw::MjVisual* ptr, py::handle owner) + : WrapperBase(ptr, owner), + headlight(&ptr_->headlight, owner_), + rgba(&ptr_->rgba, owner_) {} + +MjVisualWrapper::MjWrapper(const MjVisualWrapper& other) : MjVisualWrapper() { + *this->ptr_ = *other.ptr_; +} + +// ==================== MJMODEL ================================================ +static void MjModelCapsuleDestructor(PyObject* pyobj) { + mj_deleteModel( + static_cast(PyCapsule_GetPointer(pyobj, nullptr))); +} + +static absl::flat_hash_map& +MjModelRawPointerMap() { + static auto* hash_map = + new absl::flat_hash_map(); + return *hash_map; +} + +MjModelWrapper* MjModelWrapper::FromRawPointer(raw::MjModel* m) noexcept { + try { + auto& map = MjModelRawPointerMap(); + { + py::gil_scoped_acquire gil; + auto found = map.find(m); + return found != map.end() ? found->second : nullptr; + } + } catch (...) { + return nullptr; + } +} + +#undef MJ_M +#define MJ_M(x) ptr_->x +#define X(dtype, var, dim0, dim1) \ + , var(InitPyArray(X_ARRAY_SHAPE(ptr_->dim0, dim1), ptr_->var, owner_)) +MjModelWrapper::MjWrapper(raw::MjModel* ptr) + : WrapperBase(ptr, &MjModelCapsuleDestructor), + opt(&ptr->opt, owner_), + vis(&ptr->vis, owner_), + stat(&ptr->stat, owner_) MJMODEL_POINTERS, + text_data_bytes(ptr->text_data, ptr->ntextdata), + names_bytes(ptr->names, ptr->nnames), + paths_bytes(ptr->paths, ptr->npaths), + indexer_(ptr, owner_) { + bool is_newly_inserted = false; + { + py::gil_scoped_acquire gil; + is_newly_inserted = MjModelRawPointerMap().insert({ptr_, this}).second; + } + if (!is_newly_inserted) { + throw UnexpectedError( + "MjModelWrapper(mjModel*): MjModelRawPointerMap already contains this " + "raw mjModel*"); + } +} + +MjModelWrapper::MjWrapper(MjModelWrapper&& other) + : WrapperBase(other.ptr_, other.owner_), + opt(&ptr_->opt, owner_), + vis(&ptr_->vis, owner_), + stat(&ptr_->stat, owner_) MJMODEL_POINTERS, + text_data_bytes(ptr_->text_data, ptr_->ntextdata), + names_bytes(ptr_->names, ptr_->nnames), + paths_bytes(ptr_->paths, ptr_->npaths), + indexer_(ptr_, owner_) { + bool is_newly_inserted = false; + { + py::gil_scoped_acquire gil; + is_newly_inserted = + MjModelRawPointerMap().insert_or_assign(ptr_, this).second; + } + if (is_newly_inserted) { + throw UnexpectedError( + "MjModelRawPointerMap does not contains the moved-from mjModel*"); + } + other.ptr_ = nullptr; +} +#undef X +#undef MJ_M +#define MJ_M(x) x + +// Delegating to the MjModelWrapper::MjWrapper(raw::MjModel*) constructor, +// no need to modify MjModelRawPointerMap here. +MjModelWrapper::MjWrapper(const MjModelWrapper& other) + : MjModelWrapper(InterceptMjErrors(mj_copyModel)(NULL, other.get())) {} + +MjModelWrapper::~MjWrapper() { + if (ptr_) { + bool erased = false; + { + py::gil_scoped_acquire gil; + erased = MjModelRawPointerMap().erase(ptr_); + } + if (!erased) { + std::cerr << "MjModelRawPointerMap does not contain this raw mjModel*" + << std::endl; + std::terminate(); + } + } +} + +// Helper function for both LoadXMLFile and LoadBinaryFile. +// Creates a temporary MJB from the assets dictionary if one is supplied. +template +static raw::MjModel* LoadModelFileImpl(const std::string& filename, + const std::vector& assets, + LoadFunc&& loadfunc) { + mjVFS vfs; + mjVFS* vfs_ptr = nullptr; + if (!assets.empty()) { + mj_defaultVFS(&vfs); + vfs_ptr = &vfs; + for (const auto& asset : assets) { + std::string buffer_name = StripPath(asset.name); + const int vfs_error = InterceptMjErrors(mj_addBufferVFS)( + vfs_ptr, buffer_name.c_str(), asset.content, asset.content_size); + if (vfs_error) { + mj_deleteVFS(vfs_ptr); + if (vfs_error == 2) { + throw py::value_error("Repeated file name in assets dict: " + + buffer_name); + } else { + throw py::value_error("Asset failed to load: " + buffer_name); + } + } + } + } + + raw::MjModel* model = loadfunc(filename.c_str(), vfs_ptr); + mj_deleteVFS(vfs_ptr); + if (model && !model->buffer) { + mj_deleteModel(model); + model = nullptr; + } + return model; +} + +MjModelWrapper MjModelWrapper::LoadXMLFile( + const std::string& filename, + const std::optional>& assets) { + const auto converted_assets = ConvertAssetsDict(assets); + raw::MjModel* model; + { + py::gil_scoped_release no_gil; + char error[1024]; + model = LoadModelFileImpl(filename, converted_assets, + [&error](const char* filename, const mjVFS* vfs) { + return InterceptMjErrors(mj_loadXML)( + filename, vfs, error, sizeof(error)); + }); + if (!model) { + throw py::value_error(error); + } + } + return MjModelWrapper(model); +} + +MjModelWrapper MjModelWrapper::LoadBinaryFile( + const std::string& filename, + const std::optional>& assets) { + const auto converted_assets = ConvertAssetsDict(assets); + raw::MjModel* model; + { + py::gil_scoped_release no_gil; + model = LoadModelFileImpl(filename, converted_assets, + InterceptMjErrors(mj_loadModel)); + if (!model) { + throw py::value_error("mj_loadModel: failed to load from mjb"); + } + } + return MjModelWrapper(model); +} + +MjModelWrapper MjModelWrapper::LoadXML( + const std::string& xml, + const std::optional>& assets) { + auto converted_assets = ConvertAssetsDict(assets); + raw::MjModel* model; + { + py::gil_scoped_release no_gil; + std::string model_filename = "model_.xml"; + if (assets.has_value()) { + while (assets->find(model_filename) != assets->end()) { + model_filename = + model_filename.substr(0, model_filename.size() - 4) + "_.xml"; + } + } + converted_assets.emplace_back(model_filename.c_str(), xml.c_str(), + xml.length()); + char error[1024]; + model = LoadModelFileImpl(model_filename, converted_assets, + [&error](const char* filename, const mjVFS* vfs) { + return InterceptMjErrors(mj_loadXML)( + filename, vfs, error, sizeof(error)); + }); + if (!model) { + throw py::value_error(error); + } + } + return MjModelWrapper(model); +} + +MjModelWrapper MjModelWrapper::WrapRawModel(raw::MjModel* m) { + return MjModelWrapper(m); +} + +py::tuple RecompileSpec(raw::MjSpec* spec, const MjModelWrapper& old_m, + const MjDataWrapper& old_d) { + raw::MjModel* m = static_cast(mju_malloc(sizeof(mjModel))); + m->buffer = nullptr; + raw::MjData* d = mj_copyData(nullptr, old_m.get(), old_d.get()); + if (mj_recompile(spec, nullptr, m, d)) { + throw py::value_error(mjs_getError(spec)); + } + + py::object m_pyobj = py::cast((MjModelWrapper(m))); + py::object d_pyobj = + py::cast((MjDataWrapper(py::cast(m_pyobj), d))); + return py::make_tuple(m_pyobj, d_pyobj); +} + +namespace { +// A byte at the start of serialized mjModel structs, which can be incremented +// when we change the serialization logic to reject pickles from an unsupported +// future version. +constexpr static char kSerializationVersion = 1; + +void CheckInput(const std::istream& input, std::string class_name) { + if (input.fail()) { + throw py::value_error("Invalid serialized " + class_name + "."); + } +} + +} // namespace + +void MjModelWrapper::Serialize(std::ostream& output) const { + WriteChar(output, kSerializationVersion); + + int model_size = mj_sizeModel(get()); + WriteInt(output, model_size); + std::string buffer(model_size, 0); + mj_saveModel(get(), nullptr, buffer.data(), model_size); + WriteBytes(output, buffer.data(), model_size); +} + +std::unique_ptr MjModelWrapper::Deserialize( + std::istream& input) { + CheckInput(input, "mjModel"); + + char serializationVersion = ReadChar(input); + CheckInput(input, "mjModel"); + + if (serializationVersion != kSerializationVersion) { + throw py::value_error("Incompatible serialization version."); + } + + std::size_t model_size = ReadInt(input); + CheckInput(input, "mjModel"); + if (model_size < 0) { + throw py::value_error("Invalid serialized mjModel."); + } + std::string model_bytes(model_size, 0); + ReadBytes(input, model_bytes.data(), model_size); + CheckInput(input, "mjModel"); + + raw::MjModel* model = LoadModelFileImpl( + "model.mjb", + {{"model.mjb", model_bytes.data(), static_cast(model_size)}}, + InterceptMjErrors(mj_loadModel)); + if (!model) { + throw py::value_error("Invalid serialized mjModel."); + } + return std::unique_ptr(new MjModelWrapper(model)); +} + +// ==================== MJCONTACT ============================================== +#define X(var) var(InitPyArray(ptr_->var, owner_)) +MjContactWrapper::MjWrapper() + : WrapperBase(new raw::MjContact{}), + X(pos), + X(frame), + X(friction), + X(solref), + X(solreffriction), + X(solimp), + X(H), + X(geom), + X(flex), + X(elem), + X(vert) {} + +MjContactWrapper::MjWrapper(raw::MjContact* ptr, py::handle owner) + : WrapperBase(ptr, owner), + X(pos), + X(frame), + X(friction), + X(solref), + X(solreffriction), + X(solimp), + X(H), + X(geom), + X(flex), + X(elem), + X(vert) {} +#undef X + +MjContactWrapper::MjWrapper(const MjContactWrapper& other) + : MjContactWrapper() { + *this->ptr_ = *other.ptr_; +} + +MjContactList::MjStructList(raw::MjContact* ptr, int nconmax, int* ncon, + py::handle owner) + : StructListBase(ptr, nconmax, owner, /* lazy = */ true), ncon_(ncon) {} + +// Slicing +MjContactList::MjStructList(MjContactList& other, py::slice slice) + : StructListBase(other, slice), ncon_(other.ncon_) {} + +// ==================== MJDATA ================================================= +static void MjDataCapsuleDestructor(PyObject* pyobj) { + mj_deleteData( + static_cast(PyCapsule_GetPointer(pyobj, nullptr))); +} + +absl::flat_hash_map& MjDataRawPointerMap() { + static auto* hash_map = + new absl::flat_hash_map(); + return *hash_map; +} + +MjDataWrapper* MjDataWrapper::FromRawPointer(raw::MjData* m) noexcept { + try { + auto& map = MjDataRawPointerMap(); + { + py::gil_scoped_acquire gil; + auto found = map.find(m); + return found != map.end() ? found->second : nullptr; + } + } catch (...) { + return nullptr; + } +} + +namespace { +// default timer callback (seconds) +mjtNum GetTime() { + using Clock = std::chrono::steady_clock; + using Seconds = std::chrono::duration; + static const Clock::time_point tm_start = Clock::now(); + return Seconds(Clock::now() - tm_start).count(); +} +} // namespace + +MjDataWrapper::MjWrapper(MjModelWrapper* model) + : WrapperBase(InterceptMjErrors(mj_makeData)(model->get()), + &MjDataCapsuleDestructor), +#undef MJ_M +#define MJ_M(x) model->get()->x +#define X(dtype, var, dim0, dim1) \ + var(InitPyArray(X_ARRAY_SHAPE(model->get()->dim0, dim1), ptr_->var, owner_)), + MJDATA_POINTERS +#undef MJ_M +#define MJ_M(x) (x) +#undef X + + contact(MjContactList(ptr_->contact, NConMax(ptr_), &ptr_->ncon, owner_)), + +#define X(dtype, var, dim0, dim1) var(InitPyArray(ptr_->var, owner_)), + MJDATA_VECTOR +#undef X + model_(model), + model_ref_(py::cast(model_)), + indexer_(ptr_, model_->get(), owner_) { + bool is_newly_inserted = false; + { + py::gil_scoped_acquire gil; + is_newly_inserted = MjDataRawPointerMap().insert({ptr_, this}).second; + } + if (!is_newly_inserted) { + throw UnexpectedError( + "MjDataRawPointerMap already contains this raw mjData*"); + } + + // install default timer if not already installed + { + py::gil_scoped_acquire gil; + if (!mjcb_time) { + mjcb_time = GetTime; + } + } +} + +MjDataWrapper::MjWrapper(const MjDataWrapper& other) + : WrapperBase(other.Copy(), &MjDataCapsuleDestructor), +#undef MJ_M +#define MJ_M(x) other.model_->get()->x +#define X(dtype, var, dim0, dim1) \ + var(InitPyArray(X_ARRAY_SHAPE(other.model_->get()->dim0, dim1), ptr_->var, \ + owner_)), + MJDATA_POINTERS +#undef MJ_M +#define MJ_M(x) (x) +#undef X + + contact(MjContactList(ptr_->contact, NConMax(ptr_), &ptr_->ncon, owner_)), + +#define X(dtype, var, dim0, dim1) var(InitPyArray(ptr_->var, owner_)), + MJDATA_VECTOR +#undef X + model_(other.model_), + model_ref_(other.model_ref_), + indexer_(ptr_, model_->get(), owner_) { + bool is_newly_inserted = false; + { + py::gil_scoped_acquire gil; + is_newly_inserted = MjDataRawPointerMap().insert({ptr_, this}).second; + } + if (!is_newly_inserted) { + throw UnexpectedError( + "MjDataRawPointerMap already contains this raw mjData*"); + } +} + +MjDataWrapper::MjWrapper(MjDataWrapper&& other) + : WrapperBase(other.ptr_, other.owner_), +#undef MJ_M +#define MJ_M(x) other.model_->get()->x +#define X(dtype, var, dim0, dim1) \ + var(InitPyArray(X_ARRAY_SHAPE(other.model_->get()->dim0, dim1), ptr_->var, \ + owner_)), + MJDATA_POINTERS +#undef MJ_M +#define MJ_M(x) (x) +#undef X + + contact(MjContactList(ptr_->contact, NConMax(ptr_), &ptr_->ncon, owner_)), + +#define X(dtype, var, dim0, dim1) var(InitPyArray(ptr_->var, owner_)), + MJDATA_VECTOR +#undef X + model_(other.model_), + model_ref_(std::move(other.model_ref_)), + indexer_(ptr_, model_->get(), owner_) { + bool is_newly_inserted = false; + { + py::gil_scoped_acquire gil; + is_newly_inserted = + MjDataRawPointerMap().insert_or_assign(ptr_, this).second; + } + if (is_newly_inserted) { + throw UnexpectedError( + "MjDataRawPointerMap does not contains the moved-from mjData*"); + } + other.ptr_ = nullptr; +} + +MjDataWrapper::MjWrapper(const MjDataWrapper& other, MjModelWrapper* model) + : WrapperBase(other.Copy(), &MjDataCapsuleDestructor), +#undef MJ_M +#define MJ_M(x) other.model_->get()->x +#define X(dtype, var, dim0, dim1) \ + var(InitPyArray(X_ARRAY_SHAPE(other.model_->get()->dim0, dim1), ptr_->var, \ + owner_)), + MJDATA_POINTERS +#undef MJ_M +#define MJ_M(x) (x) +#undef X + + contact(MjContactList(ptr_->contact, NConMax(ptr_), &ptr_->ncon, owner_)), + +#define X(dtype, var, dim0, dim1) var(InitPyArray(ptr_->var, owner_)), + MJDATA_VECTOR +#undef X + model_(model), + model_ref_(py::cast(model_)), + indexer_(ptr_, model_->get(), owner_) { + bool is_newly_inserted = false; + { + py::gil_scoped_acquire gil; + is_newly_inserted = MjDataRawPointerMap().insert({ptr_, this}).second; + } + if (!is_newly_inserted) { + throw UnexpectedError( + "MjDataRawPointerMap already contains this raw mjData*"); + } +} + +MjDataWrapper::MjWrapper(MjModelWrapper* model, raw::MjData* d) + : WrapperBase(d, &MjDataCapsuleDestructor), +#undef MJ_M +#define MJ_M(x) model->get()->x +#define X(dtype, var, dim0, dim1) \ + var(InitPyArray(X_ARRAY_SHAPE(model->get()->dim0, dim1), ptr_->var, owner_)), + MJDATA_POINTERS +#undef MJ_M +#define MJ_M(x) (x) +#undef X + + contact(MjContactList(ptr_->contact, NConMax(ptr_), &ptr_->ncon, owner_)), + +#define X(dtype, var, dim0, dim1) var(InitPyArray(ptr_->var, owner_)), + MJDATA_VECTOR +#undef X + model_(model), + model_ref_(py::cast(model_)), + indexer_(ptr_, model_->get(), owner_) { + bool is_newly_inserted = false; + { + py::gil_scoped_acquire gil; + is_newly_inserted = MjDataRawPointerMap().insert({ptr_, this}).second; + } + if (!is_newly_inserted) { + throw UnexpectedError( + "MjDataRawPointerMap already contains this raw mjData*"); + } +} + +MjDataWrapper::~MjWrapper() { + if (ptr_) { + bool erased = false; + { + py::gil_scoped_acquire gil; + erased = MjDataRawPointerMap().erase(ptr_); + } + if (!erased) { + std::cerr << "MjDataRawPointerMap does not contain this raw mjData*" + << std::endl; + std::terminate(); + } + } +} + +void MjDataWrapper::Serialize(std::ostream& output) const { + // TODO: Replace this custom serialization with a protobuf + WriteChar(output, kSerializationVersion); + + model_->Serialize(output); + + // Write struct and scalar fields +#define X(var) WriteBytes(output, &ptr_->var, sizeof(ptr_->var)) + X(maxuse_stack); + X(maxuse_arena); + X(maxuse_con); + X(maxuse_efc); + X(solver); + X(timer); + X(warning); + X(ncon); + X(ne); + X(nf); + X(nJ); + X(nA); + X(nefc); + X(nisland); + X(time); + X(energy); +#undef X + + // Write buffer and arena contents + { + MJDATA_POINTERS_PREAMBLE((this->model_->get())) + +#define X(type, name, nr, nc) \ + WriteBytes(output, ptr_->name, \ + sizeof(type) * (this->model_->get()->nr) * (nc)); + MJDATA_POINTERS +#undef X + +#undef MJ_M +#define MJ_M(x) this->model_->get()->x +#undef MJ_D +#define MJ_D(x) this->ptr_->x +#define X(type, name, nr, nc) \ + if ((nr) * (nc)) { \ + WriteBytes(output, ptr_->name, \ + ptr_->name ? sizeof(type) * (nr) * (nc) : 0); \ + } + + MJDATA_ARENA_POINTERS_CONTACT + MJDATA_ARENA_POINTERS_SOLVER + if (mj_isDual(this->model_->get())) { + MJDATA_ARENA_POINTERS_DUAL + } + if (this->ptr_->nisland) { + MJDATA_ARENA_POINTERS_ISLAND + } +#undef MJ_M +#define MJ_M(x) x +#undef MJ_D +#define MJ_D(x) x +#undef X + } +} + +MjDataWrapper MjDataWrapper::Deserialize(std::istream& input) { + char serializationVersion = ReadChar(input); + CheckInput(input, "mjData"); + + if (serializationVersion != kSerializationVersion) { + throw py::value_error("Incompatible serialization version."); + } + + // Read the model that was used to create the mjData. + std::unique_ptr m_wrapper = + MjModelWrapper::Deserialize(input); + raw::MjModel& m = *m_wrapper->get(); + + bool is_dual = mj_isDual(&m); + + raw::MjData* d = mj_makeData(&m); + if (!d) { + throw py::value_error("Failed to create mjData."); + } + + // Read structs and scalar fields +#define X(var) \ + ReadBytes(input, (void*)&d->var, sizeof(d->var)); \ + CheckInput(input, "mjData"); + + X(maxuse_stack); + X(maxuse_arena); + X(maxuse_con); + X(maxuse_efc); + X(solver); + X(timer); + X(warning); + X(ncon); + X(ne); + X(nf); + X(nJ); + X(nA); + X(nefc); + X(nisland); + X(time); + X(energy); +#undef X + + // Read buffer and arena contents + { + MJDATA_POINTERS_PREAMBLE((&m)) + +#define X(type, name, nr, nc) \ + ReadBytes(input, d->name, sizeof(type) * (m.nr) * (nc)); + MJDATA_POINTERS +#undef X + +#undef MJ_M +#define MJ_M(x) m.x +#undef MJ_D +#define MJ_D(x) d->x +// arena pointers might be null, so we need to check the size before allocating. +#define X(type, name, nr, nc) \ + if ((nr) * (nc)) { \ + std::size_t actual_nbytes = ReadInt(input); \ + if (actual_nbytes) { \ + if (actual_nbytes != sizeof(type) * (nr) * (nc)) { \ + input.setstate(input.rdstate() | std::ios_base::failbit); \ + } else { \ + d->name = static_castname)>( \ + mj_arenaAllocByte(d, sizeof(type) * (nr) * (nc), alignof(type))); \ + input.read(reinterpret_cast(d->name), actual_nbytes); \ + } \ + } else { \ + d->name = nullptr; \ + } \ + } + + MJDATA_ARENA_POINTERS_CONTACT + MJDATA_ARENA_POINTERS_SOLVER + if (is_dual) { + MJDATA_ARENA_POINTERS_DUAL + } + if (d->nisland) { + MJDATA_ARENA_POINTERS_ISLAND + } +#undef MJ_M +#define MJ_M(x) x +#undef MJ_D +#define MJ_D(x) x +#undef X + } + CheckInput(input, "mjData"); + + // All bytes should have been used. + input.ignore(1); + if (!input.eof()) { + throw py::value_error("Invalid serialized mjData."); + } + + return MjDataWrapper(m_wrapper.release(), d); +} + +raw::MjData* MjDataWrapper::Copy() const { + const raw::MjModel* m = model_->get(); + return InterceptMjErrors(mj_copyData)(NULL, m, this->ptr_); +} + +// ==================== MJSTATISTIC ============================================ +#define X(var) var(InitPyArray(ptr_->var, owner_)) +MjStatisticWrapper::MjWrapper() + : WrapperBase(new raw::MjStatistic{}), X(center) {} + +MjStatisticWrapper::MjWrapper(raw::MjStatistic* ptr, py::handle owner) + : WrapperBase(ptr, owner), X(center) {} +#undef X + +MjStatisticWrapper::MjWrapper(const MjStatisticWrapper& other) + : MjStatisticWrapper() { + *this->ptr_ = *other.ptr_; +} + +// ==================== MJWARNINGSTAT ========================================== +MjWarningStatWrapper::MjWrapper() : WrapperBase(new raw::MjWarningStat{}) {} + +MjWarningStatWrapper::MjWrapper(raw::MjWarningStat* ptr, py::handle owner) + : WrapperBase(ptr, owner) {} + +MjWarningStatWrapper::MjWrapper(const MjWarningStatWrapper& other) + : MjWarningStatWrapper() { + *this->ptr_ = *other.ptr_; +} + +#define X(type, var) \ + var(std::vector{num}, std::vector{sizeof(raw::MjWarningStat)}, \ + &ptr->var, owner) +MjWarningStatList::MjStructList(raw::MjWarningStat* ptr, int num, + py::handle owner) + : StructListBase(ptr, num, owner), X(int, lastinfo), X(int, number) {} +#undef X + +// Slicing +#define X(type, var) var(other.var[slice]) +MjWarningStatList::MjStructList(MjWarningStatList& other, py::slice slice) + : StructListBase(other, slice), X(int, lastinfo), X(int, number) {} +#undef X + +// ==================== MJTIMERSTAT ============================================ +MjTimerStatWrapper::MjWrapper() : WrapperBase(new raw::MjTimerStat{}) {} + +MjTimerStatWrapper::MjWrapper(raw::MjTimerStat* ptr, py::handle owner) + : WrapperBase(ptr, owner) {} + +MjTimerStatWrapper::MjWrapper(const MjTimerStatWrapper& other) + : MjTimerStatWrapper() { + *this->ptr_ = *other.ptr_; +} + +#define X(type, var) \ + var(std::vector{num}, std::vector{sizeof(raw::MjTimerStat)}, \ + &ptr->var, owner) +MjTimerStatList::MjStructList(raw::MjTimerStat* ptr, int num, py::handle owner) + : StructListBase(ptr, num, owner), X(mjtNum, duration), X(int, number) {} +#undef X + +// Slicing +#define X(type, var) var(other.var[slice]) +MjTimerStatList::MjStructList(MjTimerStatList& other, py::slice slice) + : StructListBase(other, slice), X(mjtNum, duration), X(int, number) {} +#undef X + +// ==================== MJSOLVERSTAT =========================================== +MjSolverStatWrapper::MjWrapper() : WrapperBase(new raw::MjSolverStat{}) {} + +MjSolverStatWrapper::MjWrapper(raw::MjSolverStat* ptr, py::handle owner) + : WrapperBase(ptr, owner) {} + +MjSolverStatWrapper::MjWrapper(const MjSolverStatWrapper& other) + : MjSolverStatWrapper() { + *this->ptr_ = *other.ptr_; +} + +#define X(type, var) \ + var(std::vector{num}, std::vector{sizeof(raw::MjSolverStat)}, \ + &ptr->var, owner) +MjSolverStatList::MjStructList(raw::MjSolverStat* ptr, int num, + py::handle owner) + : StructListBase(ptr, num, owner), + X(mjtNum, improvement), + X(mjtNum, gradient), + X(mjtNum, lineslope), + X(int, nactive), + X(int, nchange), + X(int, neval), + X(int, nupdate) {} +#undef X +#undef XN + +// Slicing +#define X(type, var) var(other.var[slice]) +MjSolverStatList::MjStructList(MjSolverStatList& other, py::slice slice) + : StructListBase(other, slice), + X(mjtNum, improvement), + X(mjtNum, gradient), + X(mjtNum, lineslope), + X(int, nactive), + X(int, nchange), + X(int, neval), + X(int, nupdate) {} +#undef X + +// ==================== MJVPERTURB ============================================= +#define X(var) var(InitPyArray(ptr_->var, owner_)) +MjvPerturbWrapper::MjWrapper() + : WrapperBase([]() { + raw::MjvPerturb* const pert = new raw::MjvPerturb; + mjv_defaultPerturb(pert); + return pert; + }()), + X(refpos), + X(refquat), + X(refselpos), + X(localpos) {} +#undef X + +MjvPerturbWrapper::MjWrapper(const MjvPerturbWrapper& other) + : MjvPerturbWrapper() { + *this->ptr_ = *other.ptr_; +} + +// ==================== MJVCAMERA ============================================== +#define X(var) var(InitPyArray(ptr_->var, owner_)) +MjvCameraWrapper::MjWrapper() + : WrapperBase([]() { + raw::MjvCamera* const cam = new raw::MjvCamera; + mjv_defaultCamera(cam); + return cam; + }()), + X(lookat) {} +#undef X + +MjvCameraWrapper::MjWrapper(const MjvCameraWrapper& other) + : MjvCameraWrapper() { + *this->ptr_ = *other.ptr_; +} + +// ==================== MJVGLCAMERA ============================================ +#define X(var) var(InitPyArray(ptr_->var, owner_)) +MjvGLCameraWrapper::MjWrapper() + : WrapperBase(new raw::MjvGLCamera{}), X(pos), X(forward), X(up) {} + +MjvGLCameraWrapper::MjWrapper(raw::MjvGLCamera* ptr, py::handle owner) + : WrapperBase(ptr, owner), X(pos), X(forward), X(up) {} +#undef X + +MjvGLCameraWrapper::MjWrapper(raw::MjvGLCamera&& other) : MjvGLCameraWrapper() { + *this->ptr_ = other; +} + +MjvGLCameraWrapper::MjWrapper(const MjvGLCameraWrapper& other) + : MjvGLCameraWrapper() { + *this->ptr_ = *other.ptr_; +} + +// ==================== MJVGEOM ================================================ +#define X(var) var(InitPyArray(ptr_->var, owner_)) +MjvGeomWrapper::MjWrapper() + : WrapperBase(new raw::MjvGeom{}), + X(size), + X(pos), + mat([this]() { + static_assert(sizeof(ptr_->mat) == sizeof(ptr_->mat[0]) * 9); + return InitPyArray(std::array{3, 3}, ptr_->mat, owner_); + }()), + X(rgba) { + mjv_initGeom(ptr_, mjGEOM_NONE, nullptr, nullptr, nullptr, nullptr); +} + +MjvGeomWrapper::MjWrapper(raw::MjvGeom* ptr, py::handle owner) + : WrapperBase(ptr, owner), + X(size), + X(pos), + mat([this]() { + static_assert(sizeof(ptr_->mat) == sizeof(ptr_->mat[0]) * 9); + return InitPyArray(std::array{3, 3}, ptr_->mat, owner_); + }()), + X(rgba) {} +#undef X + +MjvGeomWrapper::MjWrapper(const MjvGeomWrapper& other) : MjvGeomWrapper() { + *this->ptr_ = *other.ptr_; +} + +// ==================== MJVLIGHT =============================================== +#define X(var) var(InitPyArray(ptr_->var, owner_)) +MjvLightWrapper::MjWrapper() + : WrapperBase(new raw::MjvLight{}), + X(pos), + X(dir), + X(attenuation), + X(ambient), + X(diffuse), + X(specular) {} + +MjvLightWrapper::MjWrapper(raw::MjvLight* ptr, py::handle owner) + : WrapperBase(ptr, owner), + X(pos), + X(dir), + X(attenuation), + X(ambient), + X(diffuse), + X(specular) {} +#undef X + +MjvLightWrapper::MjWrapper(const MjvLightWrapper& other) : MjvLightWrapper() { + *this->ptr_ = *other.ptr_; +} + +// ==================== MJVOPTION ============================================== +#define X(var) var(InitPyArray(ptr_->var, owner_)) +MjvOptionWrapper::MjWrapper() + : WrapperBase([]() { + raw::MjvOption* const opt = new raw::MjvOption; + mjv_defaultOption(opt); + return opt; + }()), + X(geomgroup), + X(sitegroup), + X(jointgroup), + X(tendongroup), + X(actuatorgroup), + X(flexgroup), + X(skingroup), + X(flags) {} +#undef X + +MjvOptionWrapper::MjWrapper(const MjvOptionWrapper& other) + : MjvOptionWrapper() { + *this->ptr_ = *other.ptr_; +} + +// ==================== MJVSCENE =============================================== +static void MjvSceneCapsuleDestructor(PyObject* pyobj) { + py::gil_scoped_acquire gil; + auto* scn = static_cast(PyCapsule_GetPointer(pyobj, nullptr)); + if (scn) { + mjv_freeScene(scn); + delete scn; + } +} + +#define X(var) var(InitPyArray(ptr_->var, owner_)) +#define XN(var, n) var(InitPyArray(std::array{n}, ptr_->var, owner_)) +MjvSceneWrapper::MjWrapper() + : WrapperBase( + []() { + raw::MjvScene* const scn = new raw::MjvScene; + mjv_defaultScene(scn); + InterceptMjErrors(mjv_makeScene)(nullptr, scn, 0); + return scn; + }(), + &MjvSceneCapsuleDestructor), + nskinvert(0), + XN(geoms, 0), + XN(geomorder, 0), + XN(flexedgeadr, 0), + XN(flexedgenum, 0), + XN(flexvertadr, 0), + XN(flexvertnum, 0), + XN(flexfaceadr, 0), + XN(flexfacenum, 0), + XN(flexfaceused, 0), + XN(flexedge, 0), + XN(flexvert, 0), + XN(flexface, 0), + XN(flexnormal, 0), + XN(flextexcoord, 0), + XN(skinfacenum, 0), + XN(skinvertadr, 0), + XN(skinvertnum, 0), + XN(skinvert, 0), + XN(skinnormal, 0), + X(lights), + X(camera), + X(translate), + X(rotate), + X(flags), + X(framergb) {} + +#define XN(var, n) var(InitPyArray(std::array{n}, ptr_->var, owner_)) +MjvSceneWrapper::MjWrapper(const MjModelWrapper& model, int maxgeom) + : WrapperBase( + [maxgeom](const raw::MjModel* m) { + raw::MjvScene* const scn = new raw::MjvScene; + mjv_defaultScene(scn); + InterceptMjErrors(mjv_makeScene)(m, scn, maxgeom); + return scn; + }(model.get()), + &MjvSceneCapsuleDestructor), + nskinvert([](const raw::MjModel* m) { + int nskinvert = 0; + for (int i = 0; i < m->nskin; ++i) { + nskinvert += m->skin_vertnum[i]; + } + return nskinvert; + }(model.get())), + nflexface([](const raw::MjModel* m) { + int nflexface = 0; + int flexfacenum = 0; + for (int f = 0; f < m->nflex; f++) { + if (m->flex_dim[f] == 0) { + // 1D : 0 + flexfacenum = 0; + } else if (m->flex_dim[f] == 2) { + // 2D: 2*fragments + 2*elements + flexfacenum = 2 * m->flex_shellnum[f] + 2 * m->flex_elemnum[f]; + } else { + // 3D: max(fragments, 4*maxlayer) + // find number of elements in biggest layer + int maxlayer = 0, layer = 0, nlayer = 1; + while (nlayer) { + nlayer = 0; + for (int e = 0; e < m->flex_elemnum[f]; e++) { + if (m->flex_elemlayer[m->flex_elemadr[f] + e] == layer) { + nlayer++; + } + } + maxlayer = mjMAX(maxlayer, nlayer); + layer++; + } + flexfacenum = mjMAX(m->flex_shellnum[f], 4 * maxlayer); + } + + // accumulate over flexes + nflexface += flexfacenum; + } + return nflexface; + }(model.get())), + nflexedge(model.get()->nflexedge), + nflexvert(model.get()->nflexvert), + XN(geoms, ptr_->maxgeom), + XN(geomorder, ptr_->maxgeom), + XN(flexedgeadr, ptr_->nflex), + XN(flexedgenum, ptr_->nflex), + XN(flexvertadr, ptr_->nflex), + XN(flexvertnum, ptr_->nflex), + XN(flexfaceadr, ptr_->nflex), + XN(flexfacenum, ptr_->nflex), + XN(flexfaceused, ptr_->nflex), + XN(flexedge, 2 * nflexedge), + XN(flexvert, 3 * nflexvert), + XN(flexface, 9 * nflexface), + XN(flexnormal, 9 * nflexface), + XN(flextexcoord, 6 * nflexface), + XN(skinfacenum, ptr_->nskin), + XN(skinvertadr, ptr_->nskin), + XN(skinvertnum, ptr_->nskin), + XN(skinvert, 3 * nskinvert), + XN(skinnormal, 3 * nskinvert), + X(lights), + X(camera), + X(translate), + X(rotate), + X(flags), + X(framergb) {} +#undef X +#undef XN + +template +static T* MallocAndCopy(const T* src, int count) { + if (src) { + T* out = static_cast(mju_malloc(count * sizeof(T))); + std::memcpy(out, src, count * sizeof(T)); + return out; + } else { + return nullptr; + } +} + +MjvSceneWrapper::MjWrapper(const MjvSceneWrapper& other) : MjvSceneWrapper() { + mjv_freeScene(ptr_); + *ptr_ = *other.ptr_; + +#define XN(var, n) \ + ptr_->var = MallocAndCopy(other.ptr_->var, n); \ + var = InitPyArray(std::array{n}, ptr_->var, owner_); + + XN(geoms, ptr_->ngeom); + XN(geomorder, ptr_->ngeom); + XN(flexedgeadr, ptr_->nflex); + XN(flexedgenum, ptr_->nflex); + XN(flexvertadr, ptr_->nflex); + XN(flexvertnum, ptr_->nflex); + XN(flexfaceadr, ptr_->nflex); + XN(flexfacenum, ptr_->nflex); + XN(flexfaceused, ptr_->nflex); + XN(flexedge, 2 * nflexedge); + XN(flexvert, 3 * nflexvert); + XN(flexface, 9 * nflexface); + XN(flexnormal, 9 * nflexface); + XN(flextexcoord, 6 * nflexface); + XN(skinfacenum, ptr_->nskin); + XN(skinvertadr, ptr_->nskin); + XN(skinvertnum, ptr_->nskin); + XN(skinvert, 3 * nskinvert); + XN(skinnormal, 3 * nskinvert); + +#undef XN +} + +// ==================== MJVFIGURE ============================================== +#define X(var) var(InitPyArray(ptr_->var, owner_)) +MjvFigureWrapper::MjWrapper() + : WrapperBase([]() { + raw::MjvFigure* const fig = new raw::MjvFigure; + mjv_defaultFigure(fig); + return fig; + }()), + X(flg_ticklabel), + X(gridsize), + X(gridrgb), + X(figurergba), + X(panergba), + X(legendrgba), + X(textrgb), + X(linergb), + X(range), + X(highlight), + X(linepnt), + X(linedata), + X(xaxispixel), + X(yaxispixel), + X(xaxisdata), + X(yaxisdata), +#undef X + + linename([](raw::MjvFigure* ptr, py::handle owner) { +// Use a macro to help us static_assert that the array extents here are kept +// in sync with mjVisualize.h. +#define MAKE_STR_ARRAY(N1, N2) \ + static_assert( \ + std::is_same_v); \ + return py::array(py::dtype("|S" #N2), N1, ptr->linename, owner); + MAKE_STR_ARRAY(mjMAXLINE, 100); + +#undef MAKE_STR_ARRAY + }(ptr_, owner_)) { +} + +MjvFigureWrapper::MjWrapper(const MjvFigureWrapper& other) + : MjvFigureWrapper() { + *this->ptr_ = *other.ptr_; +} + +} // namespace mujoco::python::_impl