Copy plugins when attaching a model.

Before this CL, the same mjCPlugin objects were referenced by both models.

Fixes #2100.

PiperOrigin-RevId: 680579479
Change-Id: Ida3459f78f9fa4a942cea31b3d5d3e8da2832feb
This commit is contained in:
Alessio Quaglino
2024-09-30 08:06:54 -07:00
committed by Copybara-Service
parent 43a82e08be
commit e6e6e3b51b
3 changed files with 83 additions and 14 deletions
+33 -13
View File
@@ -24,11 +24,11 @@
#include <cstdlib>
#include <cstring>
#include <exception>
#include <map>
#include <mutex>
#include <string>
#include <string_view>
#include <thread>
#include <unordered_map>
#include <vector>
#include <mujoco/mjdata.h>
@@ -164,9 +164,6 @@ mjCModel::mjCModel() {
// point to model from spec
PointToLocal();
// this class allocated the plugins
plugin_owner = true;
}
@@ -180,7 +177,6 @@ mjCModel::mjCModel(const mjCModel& other) {
mjCModel& mjCModel::operator=(const mjCModel& other) {
if (this != &other) {
plugin_owner = false;
this->spec = other.spec;
*static_cast<mjCModel_*>(this) = static_cast<const mjCModel_&>(other);
*static_cast<mjSpec*>(this) = static_cast<const mjSpec&>(other);
@@ -302,6 +298,22 @@ void mjCModel::SaveDofOffsets() {
template <class T>
static void mapplugin(
const std::unordered_map<mjCPlugin*, mjCPlugin*>& plugin_map, std::vector<T*>& list) {
for (const auto& element : list) {
if (element->spec.plugin.element) {
mjCPlugin* plugin = static_cast<mjCPlugin*>(element->spec.plugin.element);
// the referenced plugin might already exist in the source model
if (plugin_map.find(plugin) != plugin_map.end()) {
element->spec.plugin.element = plugin_map.at(plugin);
}
}
}
}
mjCModel& mjCModel::operator+=(const mjCModel& other) {
// create global lists
mjCBody *world = bodies_[0];
@@ -323,6 +335,21 @@ mjCModel& mjCModel::operator+=(const mjCModel& other) {
for (const auto& key : other.key_pending_) {
key_pending_.push_back(key);
}
// create new plugins and map them
std::unordered_map<mjCPlugin*, mjCPlugin*> plugin_map;
for (const auto& plugin : other.plugins_) {
plugins_.push_back(new mjCPlugin(*plugin));
plugin_map[plugin] = plugins_.back();
}
mapplugin(plugin_map, bodies_);
mapplugin(plugin_map, geoms_);
mapplugin(plugin_map, meshes_);
mapplugin(plugin_map, actuators_);
mapplugin(plugin_map, sensors_);
for (const auto& active_plugin : other.active_plugins_) {
active_plugins_.emplace_back(active_plugin);
}
}
CopyList(flexes_, other.flexes_);
CopyList(pairs_, other.pairs_);
@@ -335,10 +362,6 @@ mjCModel& mjCModel::operator+=(const mjCModel& other) {
CopyList(texts_, other.texts_);
CopyList(tuples_, other.tuples_);
// plugins are global
plugins_ = other.plugins_;
active_plugins_ = other.active_plugins_;
// restore to the original state
if (!compiled) {
ResetTreeLists();
@@ -613,10 +636,7 @@ mjCModel::~mjCModel() {
for (int i=0; i<keys_.size(); i++) delete keys_[i];
for (int i=0; i<defaults_.size(); i++) delete defaults_[i];
for (int i=0; i<specs_.size(); i++) mj_deleteSpec(specs_[i]);
if (plugin_owner) {
for (int i=0; i<plugins_.size(); i++) delete plugins_[i];
}
for (int i=0; i<plugins_.size(); i++) delete plugins_[i];
// clear sizes and pointer lists created in Compile
Clear();
-1
View File
@@ -386,7 +386,6 @@ class mjCModel : public mjCModel_, private mjSpec {
mjListKeyMap ids; // map from object names to ids
mjCError errInfo; // last error info
bool plugin_owner; // this class allocated the plugins
std::vector<mjKeyInfo> key_pending_; // attached keyframes
};
#endif // MUJOCO_SRC_USER_USER_MODEL_H_
+50
View File
@@ -216,6 +216,56 @@ TEST_F(PluginTest, DeletePlugin) {
mj_deleteModel(newmodel);
}
TEST_F(PluginTest, AttachPlugin) {
static constexpr char xml_1[] = R"(
<mujoco model="MuJoCo Model">
<worldbody>
<body name="body"/>
</worldbody>
</mujoco>)";
static constexpr char xml_2[] = R"(
<mujoco model="MuJoCo Model">
<extension>
<plugin plugin="mujoco.pid">
<instance name="actuator-1">
<config key="ki" value="4.0"/>
<config key="slewmax" value="3.14159"/>
</instance>
</plugin>
</extension>
<worldbody>
<body name="body">
<joint name="joint"/>
<geom size="0.1"/>
</body>
</worldbody>
<actuator>
<plugin name="actuator-1" plugin="mujoco.pid" instance="actuator-1"
joint="joint" actdim="2"/>
</actuator>
</mujoco>)";
std::array<char, 1000> err;
mjSpec* spec_1 = mj_parseXMLString(xml_1, 0, err.data(), err.size());
ASSERT_THAT(spec_1, NotNull()) << err.data();
mjSpec* spec_2 = mj_parseXMLString(xml_2, 0, err.data(), err.size());
ASSERT_THAT(spec_2, NotNull()) << err.data();
mjsBody* body_1 = mjs_findBody(spec_1, "body");
EXPECT_THAT(body_1, NotNull());
mjsFrame* attachment_frame = mjs_addFrame(body_1, 0);
EXPECT_THAT(attachment_frame, NotNull());
mjs_attachBody(attachment_frame, mjs_findBody(spec_2, "body"), "child-", "");
mjModel* model = mj_compile(spec_1, nullptr);
EXPECT_THAT(model, NotNull());
mj_deleteModel(model);
mj_deleteSpec(spec_1);
mj_deleteSpec(spec_2);
}
TEST_F(MujocoTest, RecompileFails) {
mjSpec* spec = mj_makeSpec();
mjsBody* body = mjs_addBody(mjs_findBody(spec, "world"), 0);