Internal Python bindings cleanup.

PiperOrigin-RevId: 495868304
Change-Id: Ib9cfe5f7a1d49126659176207bd6b4a314f8516a
This commit is contained in:
Saran Tunyasuvunakool
2022-12-16 07:22:01 -08:00
committed by Copybara-Service
parent 189349f421
commit 725b45b577
3 changed files with 23 additions and 19 deletions
+20 -17
View File
@@ -15,6 +15,9 @@
#ifndef MUJOCO_PYTHON_MJDATA_META_H_
#define MUJOCO_PYTHON_MJDATA_META_H_
#include <cstring>
#include <memory>
#include <mujoco/mujoco.h>
#include <mujoco/mjxmacro.h>
#include "raw.h"
@@ -67,6 +70,23 @@ template <typename T> class MjWrapper;
struct MjDataMetadata {
public:
friend class _impl::MjWrapper<raw::MjData>;
explicit MjDataMetadata(const raw::MjModel* m)
:
#define X(var) var(m->var),
MJMODEL_INTS
#undef X
#define X(dtype, var, n) \
var([](dtype* src, int len) { \
dtype* dst = new dtype[len]; \
std::memcpy(dst, src, len * sizeof(dtype)); \
return dst; \
}(m->var, m->n)),
MJDATA_METADATA
#undef X
is_dual(mj_isDual(m)) {
}
#define X(var) int var;
MJMODEL_INTS
@@ -81,23 +101,6 @@ struct MjDataMetadata {
MjDataMetadata() = default;
MjDataMetadata(const MjDataMetadata& other) = default;
MjDataMetadata(MjDataMetadata&& other) = default;
explicit MjDataMetadata(const raw::MjModel* m)
:
#define X(var) var(m->var),
MJMODEL_INTS
#undef X
#define X(dtype, var, n) \
var( \
[](dtype* src, int len) { \
dtype* dst = new dtype[len]; \
std::memcpy(dst, src, len * sizeof(dtype)); \
return dst; \
}(m->var, m->n)),
MJDATA_METADATA
#undef X
is_dual(mj_isDual(m)) {}
};
} // namespace mujoco::python
+1
View File
@@ -39,6 +39,7 @@
#include "errors.h"
#include "function_traits.h"
#include "indexers.h"
#include "mjdata_meta.h"
#include "raw.h"
#include "serialization.h"
#include <pybind11/numpy.h>
+2 -2
View File
@@ -477,7 +477,7 @@ class MjWrapper<raw::MjModel> : public WrapperBase<raw::MjModel> {
pybind11::bytes text_data_bytes;
pybind11::bytes names_bytes;
private:
protected:
explicit MjWrapper(raw::MjModel* ptr);
MjModelIndexer indexer_;
@@ -584,7 +584,7 @@ class MjWrapper<raw::MjData>: public WrapperBase<raw::MjData> {
py_array_or_tuple_t<mjtNum> solver_fwdinv;
py_array_or_tuple_t<mjtNum> energy;
private:
protected:
// Internal constructor which takes ownership of given mjData pointer.
// Used for deserialization.
explicit MjWrapper(MjDataMetadata&& metadata, raw::MjData* d);