Python bindings: Store a reference to MjModel in MjDataWrapper.
Before this change, MjDataWrapper contained a custom data structure called MjDataMetadata, which contained the sizes of various arrays needed to create MjData. This was used for serialization and deserialization, as well as copying MjData instances. This worked fine for native MuJoCo models, but as soon as plugins were used, there was information in the model that was needed and not available in MjDataMetadata. Since plugins are so general, keeping MjDataMetadata was untenable. PiperOrigin-RevId: 578205355 Change-Id: I11b9ee797d1da5aaf1fecbbe170b3b0f1f1a3e6c
This commit is contained in:
committed by
Copybara-Service
parent
64a59bb2dc
commit
084facc9ab
+82
-75
@@ -37,10 +37,11 @@
|
||||
#include "errors.h"
|
||||
#include "function_traits.h"
|
||||
#include "indexers.h"
|
||||
#include "mjdata_meta.h"
|
||||
#include "private.h"
|
||||
#include "raw.h"
|
||||
#include "serialization.h"
|
||||
#include <pybind11/cast.h>
|
||||
#include <pybind11/detail/common.h>
|
||||
#include <pybind11/numpy.h>
|
||||
#include <pybind11/operators.h>
|
||||
#include <pybind11/pybind11.h>
|
||||
@@ -453,7 +454,8 @@ void MjModelWrapper::Serialize(std::ostream& output) const {
|
||||
WriteBytes(output, buffer.data(), model_size);
|
||||
}
|
||||
|
||||
MjModelWrapper MjModelWrapper::Deserialize(std::istream& input) {
|
||||
std::unique_ptr<MjModelWrapper> MjModelWrapper::Deserialize(
|
||||
std::istream& input) {
|
||||
CheckInput(input, "mjModel");
|
||||
|
||||
char serializationVersion = ReadChar(input);
|
||||
@@ -479,7 +481,7 @@ MjModelWrapper MjModelWrapper::Deserialize(std::istream& input) {
|
||||
if (!model) {
|
||||
throw py::value_error("Invalid serialized mjModel.");
|
||||
}
|
||||
return MjModelWrapper(model);
|
||||
return std::unique_ptr<MjModelWrapper>(new MjModelWrapper(model));
|
||||
}
|
||||
|
||||
// ==================== MJCONTACT ==============================================
|
||||
@@ -554,13 +556,13 @@ MjDataWrapper* MjDataWrapper::FromRawPointer(raw::MjData* m) noexcept {
|
||||
}
|
||||
}
|
||||
|
||||
MjDataWrapper::MjWrapper(const MjModelWrapper& model)
|
||||
: WrapperBase(InterceptMjErrors(mj_makeData)(model.get()),
|
||||
MjDataWrapper::MjWrapper(MjModelWrapper* model)
|
||||
: WrapperBase(InterceptMjErrors(mj_makeData)(model->get()),
|
||||
&MjDataCapsuleDestructor),
|
||||
#undef MJ_M
|
||||
#define MJ_M(x) model.get()->x
|
||||
#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_)),
|
||||
var(InitPyArray(X_ARRAY_SHAPE(model->get()->dim0, dim1), ptr_->var, owner_)),
|
||||
MJDATA_POINTERS
|
||||
#undef MJ_M
|
||||
#define MJ_M(x) (x)
|
||||
@@ -571,8 +573,9 @@ MjDataWrapper::MjWrapper(const MjModelWrapper& model)
|
||||
#define X(dtype, var, dim0, dim1) var(InitPyArray(ptr_->var, owner_)),
|
||||
MJDATA_VECTOR
|
||||
#undef X
|
||||
metadata_(model.get()),
|
||||
indexer_(ptr_, &metadata_, owner_) {
|
||||
model_(model),
|
||||
model_ref_(py::cast(model_)),
|
||||
indexer_(ptr_, model_->get(), owner_) {
|
||||
bool is_newly_inserted = false;
|
||||
{
|
||||
py::gil_scoped_acquire gil;
|
||||
@@ -587,9 +590,9 @@ MjDataWrapper::MjWrapper(const MjModelWrapper& model)
|
||||
MjDataWrapper::MjWrapper(const MjDataWrapper& other)
|
||||
: WrapperBase(other.Copy(), &MjDataCapsuleDestructor),
|
||||
#undef MJ_M
|
||||
#define MJ_M(x) other.metadata_.x
|
||||
#define MJ_M(x) other.model_->get()->x
|
||||
#define X(dtype, var, dim0, dim1) \
|
||||
var(InitPyArray(X_ARRAY_SHAPE(other.metadata_.dim0, dim1), ptr_->var, \
|
||||
var(InitPyArray(X_ARRAY_SHAPE(other.model_->get()->dim0, dim1), ptr_->var, \
|
||||
owner_)),
|
||||
MJDATA_POINTERS
|
||||
#undef MJ_M
|
||||
@@ -601,8 +604,9 @@ MjDataWrapper::MjWrapper(const MjDataWrapper& other)
|
||||
#define X(dtype, var, dim0, dim1) var(InitPyArray(ptr_->var, owner_)),
|
||||
MJDATA_VECTOR
|
||||
#undef X
|
||||
metadata_(other.metadata_),
|
||||
indexer_(ptr_, &metadata_, owner_) {
|
||||
model_(other.model_),
|
||||
model_ref_(other.model_ref_),
|
||||
indexer_(ptr_, model_->get(), owner_) {
|
||||
bool is_newly_inserted = false;
|
||||
{
|
||||
py::gil_scoped_acquire gil;
|
||||
@@ -617,9 +621,9 @@ MjDataWrapper::MjWrapper(const MjDataWrapper& other)
|
||||
MjDataWrapper::MjWrapper(MjDataWrapper&& other)
|
||||
: WrapperBase(other.ptr_, other.owner_),
|
||||
#undef MJ_M
|
||||
#define MJ_M(x) other.metadata_.x
|
||||
#define MJ_M(x) other.model_->get()->x
|
||||
#define X(dtype, var, dim0, dim1) \
|
||||
var(InitPyArray(X_ARRAY_SHAPE(other.metadata_.dim0, dim1), ptr_->var, \
|
||||
var(InitPyArray(X_ARRAY_SHAPE(other.model_->get()->dim0, dim1), ptr_->var, \
|
||||
owner_)),
|
||||
MJDATA_POINTERS
|
||||
#undef MJ_M
|
||||
@@ -631,8 +635,9 @@ MjDataWrapper::MjWrapper(MjDataWrapper&& other)
|
||||
#define X(dtype, var, dim0, dim1) var(InitPyArray(ptr_->var, owner_)),
|
||||
MJDATA_VECTOR
|
||||
#undef X
|
||||
metadata_(other.metadata_),
|
||||
indexer_(ptr_, &metadata_, owner_) {
|
||||
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;
|
||||
@@ -646,12 +651,13 @@ MjDataWrapper::MjWrapper(MjDataWrapper&& other)
|
||||
other.ptr_ = nullptr;
|
||||
}
|
||||
|
||||
MjDataWrapper::MjWrapper(MjDataMetadata&& metadata, raw::MjData* d)
|
||||
: WrapperBase(d, &MjDataCapsuleDestructor),
|
||||
MjDataWrapper::MjWrapper(const MjDataWrapper& other, MjModelWrapper* model)
|
||||
: WrapperBase(other.Copy(), &MjDataCapsuleDestructor),
|
||||
#undef MJ_M
|
||||
#define MJ_M(x) metadata.x
|
||||
#define X(dtype, var, dim0, dim1) \
|
||||
var(InitPyArray(X_ARRAY_SHAPE(metadata.dim0, dim1), ptr_->var, owner_)),
|
||||
#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)
|
||||
@@ -662,8 +668,39 @@ MjDataWrapper::MjWrapper(MjDataMetadata&& metadata, raw::MjData* d)
|
||||
#define X(dtype, var, dim0, dim1) var(InitPyArray(ptr_->var, owner_)),
|
||||
MJDATA_VECTOR
|
||||
#undef X
|
||||
metadata_(std::move(metadata)),
|
||||
indexer_(ptr_, &metadata_, owner_) {
|
||||
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;
|
||||
@@ -694,19 +731,7 @@ void MjDataWrapper::Serialize(std::ostream& output) const {
|
||||
// TODO: Replace this custom serialization with a protobuf
|
||||
WriteChar(output, kSerializationVersion);
|
||||
|
||||
// Write all size fields
|
||||
#define X(var) WriteInt(output, this->metadata_.var);
|
||||
MJMODEL_INTS
|
||||
#undef X
|
||||
|
||||
WriteInt(output, this->metadata_.is_dual);
|
||||
|
||||
#define X(dtype, var, n) \
|
||||
WriteBytes(output, this->metadata_.var.get(), \
|
||||
this->metadata_.n * sizeof(dtype));
|
||||
|
||||
MJDATA_METADATA
|
||||
#undef X
|
||||
model_->Serialize(output);
|
||||
|
||||
// Write struct and scalar fields
|
||||
#define X(var) WriteBytes(output, &ptr_->var, sizeof(ptr_->var))
|
||||
@@ -729,15 +754,15 @@ void MjDataWrapper::Serialize(std::ostream& output) const {
|
||||
|
||||
// Write buffer contents
|
||||
{
|
||||
MJDATA_POINTERS_PREAMBLE((&this->metadata_))
|
||||
MJDATA_POINTERS_PREAMBLE((this->model_->get()))
|
||||
|
||||
#define X(type, name, nr, nc) \
|
||||
WriteBytes(output, ptr_->name, sizeof(type)*(this->metadata_.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->metadata_.x
|
||||
#define MJ_M(x) this->model_->get()->x
|
||||
#undef MJ_D
|
||||
#define MJ_D(x) this->ptr_->x
|
||||
#define X(type, name, nr, nc) \
|
||||
@@ -747,7 +772,7 @@ void MjDataWrapper::Serialize(std::ostream& output) const {
|
||||
|
||||
MJDATA_ARENA_POINTERS_CONTACT
|
||||
MJDATA_ARENA_POINTERS_PRIMAL
|
||||
if (this->metadata_.is_dual) {
|
||||
if (mj_isDual(this->model_->get())) {
|
||||
MJDATA_ARENA_POINTERS_DUAL
|
||||
}
|
||||
if (this->ptr_->nisland) {
|
||||
@@ -769,28 +794,12 @@ MjDataWrapper MjDataWrapper::Deserialize(std::istream& input) {
|
||||
throw py::value_error("Incompatible serialization version.");
|
||||
}
|
||||
|
||||
// Read all size and address fields
|
||||
MjDataMetadata metadata;
|
||||
raw::MjModel m{0};
|
||||
// Read the model that was used to create the mjData.
|
||||
std::unique_ptr<MjModelWrapper> m_wrapper =
|
||||
MjModelWrapper::Deserialize(input);
|
||||
raw::MjModel& m = *m_wrapper->get();
|
||||
|
||||
#define X(var) \
|
||||
metadata.var = ReadInt(input); \
|
||||
CheckInput(input, "mjData"); \
|
||||
m.var = metadata.var;
|
||||
|
||||
MJMODEL_INTS
|
||||
#undef X
|
||||
|
||||
metadata.is_dual = ReadInt(input);
|
||||
|
||||
#define X(dtype, var, n) \
|
||||
metadata.var.reset(new dtype[metadata.n]); \
|
||||
ReadBytes(input, metadata.var.get(), metadata.n * sizeof(dtype)); \
|
||||
CheckInput(input, "mjData"); \
|
||||
m.var = metadata.var.get();
|
||||
|
||||
MJDATA_METADATA
|
||||
#undef X
|
||||
bool is_dual = mj_isDual(&m);
|
||||
|
||||
raw::MjData* d = mj_makeData(&m);
|
||||
if (!d) {
|
||||
@@ -841,7 +850,7 @@ MjDataWrapper MjDataWrapper::Deserialize(std::istream& input) {
|
||||
|
||||
MJDATA_ARENA_POINTERS_CONTACT
|
||||
MJDATA_ARENA_POINTERS_PRIMAL
|
||||
if (metadata.is_dual) {
|
||||
if (is_dual) {
|
||||
MJDATA_ARENA_POINTERS_DUAL
|
||||
}
|
||||
if (d->nisland) {
|
||||
@@ -861,18 +870,12 @@ MjDataWrapper MjDataWrapper::Deserialize(std::istream& input) {
|
||||
throw py::value_error("Invalid serialized mjData.");
|
||||
}
|
||||
|
||||
return MjDataWrapper(std::move(metadata), d);
|
||||
return MjDataWrapper(m_wrapper.release(), d);
|
||||
}
|
||||
|
||||
raw::MjData* MjDataWrapper::Copy() const {
|
||||
raw::MjModel m{0};
|
||||
#define X(var) m.var = this->metadata_.var;
|
||||
MJMODEL_INTS
|
||||
#undef X
|
||||
#define X(dtype, var, n) m.var = this->metadata_.var.get();
|
||||
MJDATA_METADATA
|
||||
#undef X
|
||||
return InterceptMjErrors(mj_copyData)(NULL, &m, this->ptr_);
|
||||
const raw::MjModel* m = model_->get();
|
||||
return InterceptMjErrors(mj_copyData)(NULL, m, this->ptr_);
|
||||
}
|
||||
|
||||
// ==================== MJSTATISTIC ============================================
|
||||
@@ -1890,15 +1893,19 @@ This is useful for example when the MJB is not available as a file on disk.)"));
|
||||
|
||||
// ==================== MJDATA ===============================================
|
||||
py::class_<MjDataWrapper> mjData(m, "MjData");
|
||||
mjData.def(py::init<const MjModelWrapper&>());
|
||||
mjData.def(py::init<MjModelWrapper*>());
|
||||
mjData.def_property_readonly("_address", [](const MjDataWrapper& d) {
|
||||
return reinterpret_cast<std::uintptr_t>(d.get());
|
||||
});
|
||||
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) {
|
||||
return MjDataWrapper(other);
|
||||
MjModelWrapper* model_copy = new MjModelWrapper(other.model());
|
||||
return MjDataWrapper(other, model_copy);
|
||||
});
|
||||
mjData.def(py::pickle(
|
||||
[](const MjDataWrapper& d) { // __getstate__
|
||||
@@ -1927,7 +1934,7 @@ This is useful for example when the MJB is not available as a file on disk.)"));
|
||||
#undef X
|
||||
|
||||
#undef MJ_M
|
||||
#define MJ_M(x) d.metadata().x
|
||||
#define MJ_M(x) d.model().get()->x
|
||||
#undef MJ_D
|
||||
#define MJ_D(x) d.get()->x
|
||||
#define X(dtype, var, dim0, dim1) \
|
||||
|
||||
Reference in New Issue
Block a user