From 01d4c4675369359f672ba4ed98929e773c474840 Mon Sep 17 00:00:00 2001 From: Alessio Quaglino Date: Fri, 28 Mar 2025 06:26:30 -0700 Subject: [PATCH] Store the `mjsCompiler` -> appended `mjSpec` map when appending an `mjSpec`. Previously, we stored the source `mjSpec` during a copy as a hack for having access to the compiler options, but this is not robust since we cannot guarantee that 1) the source `mjSpec` is not destroyed before we need to look up the compiler options nor 2) that the `mjSpec` was appended without a copy. While 2) could be solved by simply handling an additional case in `mjCModel::FindSpec`, using a map also solves 1) and it is easier to understand. PiperOrigin-RevId: 741504264 Change-Id: Iab1bfd9e61299a94fa8caf3a244c067b09d54384 --- src/user/user_model.cc | 29 +++++++++++------------ src/user/user_model.h | 13 ++++------- src/user/user_objects.cc | 4 ++-- test/user/user_api_test.cc | 48 ++++++++++++++++++++++++++++++++++++++ 4 files changed, 69 insertions(+), 25 deletions(-) diff --git a/src/user/user_model.cc b/src/user/user_model.cc index a594538a..a63f13c9 100644 --- a/src/user/user_model.cc +++ b/src/user/user_model.cc @@ -212,9 +212,6 @@ mjCModel::mjCModel() { // create mjCBase lists from children lists CreateObjectLists(); - // the source spec is the model itself, overwritten in the copy constructor - source_spec_ = &spec; - // set the signature spec.element->signature = 0; } @@ -223,7 +220,6 @@ mjCModel::mjCModel() { mjCModel::mjCModel(const mjCModel& other) { CreateObjectLists(); - source_spec_ = (mjSpec*)&other.spec; *this = other; } @@ -239,6 +235,7 @@ mjCModel& mjCModel::operator=(const mjCModel& other) { // copy attached specs first so that we can resolve references to them for (const auto* s : other.specs_) { specs_.push_back(mj_copySpec(s)); + compiler2spec_[&s->compiler] = specs_.back(); } // the world copy constructor takes care of copying the tree @@ -1165,9 +1162,13 @@ mjCPlugin* mjCModel::AddPlugin() { // append spec to spec -void mjCModel::AppendSpec(mjSpec* spec) { +void mjCModel::AppendSpec(mjSpec* spec, const mjsCompiler* compiler_) { // TODO: check if the spec is already in the list specs_.push_back(spec); + + if (compiler_) { + compiler2spec_[compiler_] = spec; + } } @@ -1461,10 +1462,15 @@ mjSpec* mjCModel::FindSpec(std::string name) const { // find spec by mjsCompiler pointer -mjSpec* mjCModel::FindSpec(const mjsCompiler* compiler_) const { - if (&GetSourceSpec()->compiler == compiler_) { - return (mjSpec*)&spec; +mjSpec* mjCModel::FindSpec(const mjsCompiler* compiler_) { + if (compiler_ == &spec.compiler) { + return &spec; } + + if (compiler2spec_.find(compiler_) != compiler2spec_.end()) { + return compiler2spec_[compiler_]; + } + for (auto s : specs_) { mjSpec* source = static_cast(s->element)->FindSpec(compiler_); if (source) { @@ -1476,13 +1482,6 @@ mjSpec* mjCModel::FindSpec(const mjsCompiler* compiler_) const { -// get the spec from which this model was created -mjSpec* mjCModel::GetSourceSpec() const { - return source_spec_; -} - - - //------------------------------- COMPILER PHASES -------------------------------------------------- // make lists of objects in tree: bodies, geoms, joints, sites, cameras, lights diff --git a/src/user/user_model.h b/src/user/user_model.h index 81b94752..c55384bf 100644 --- a/src/user/user_model.h +++ b/src/user/user_model.h @@ -215,7 +215,9 @@ class mjCModel : public mjCModel_, private mjSpec { mjCTuple* AddTuple(); mjCKey* AddKey(); mjCPlugin* AddPlugin(); - void AppendSpec(mjSpec* spec); + + // append spec to this model, optionally map compiler options to the appended spec + void AppendSpec(mjSpec* spec, const mjsCompiler* compiler = nullptr); // delete elements marked as discard=true template void Delete(std::vector& elements, @@ -248,7 +250,7 @@ class mjCModel : public mjCModel_, private mjSpec { mjCBase* FindObject(mjtObj type, std::string name) const; // find object given type and name mjCBase* FindTree(mjCBody* body, mjtObj type, std::string name); // find tree object given name mjSpec* FindSpec(std::string name) const; // find spec given name - mjSpec* FindSpec(const mjsCompiler* compiler_) const; // find spec given mjsCompiler + mjSpec* FindSpec(const mjsCompiler* compiler_); // find spec given mjsCompiler void ActivatePlugin(const mjpPlugin* plugin, int slot); // activate plugin // accessors @@ -316,9 +318,6 @@ class mjCModel : public mjCModel_, private mjSpec { // map from default class name to default class pointer std::unordered_map def_map; - // get the spec from which this model was created - mjSpec* GetSourceSpec() const; - // set deepcopy flag void SetDeepCopy(bool deepcopy) { deepcopy_ = deepcopy; } @@ -332,9 +331,6 @@ class mjCModel : public mjCModel_, private mjSpec { // settings for each defaults class std::vector defaults_; - // spec from which this model was created in copy constructor - mjSpec* source_spec_; - // list of active plugins std::vector> active_plugins_; @@ -453,5 +449,6 @@ class mjCModel : public mjCModel_, private mjSpec { bool deepcopy_; // copy objects when attaching bool attached_ = false; // true if model is attached to a parent model int uid_count_ = 0; // unique id count for all objects + std::unordered_map compiler2spec_; // map from compiler to spec }; #endif // MUJOCO_SRC_USER_USER_MODEL_H_ diff --git a/src/user/user_objects.cc b/src/user/user_objects.cc index 5154af48..1e9cf407 100644 --- a/src/user/user_objects.cc +++ b/src/user/user_objects.cc @@ -904,7 +904,7 @@ mjCBody& mjCBody::operator+=(const mjCBody& other) { mjCBody& mjCBody::operator+=(const mjCFrame& other) { // append a copy of the attached spec if (other.model != model && !model->FindSpec(mjs_getString(other.model->spec.modelname))) { - model->AppendSpec(mj_copySpec(&other.model->spec)); + model->AppendSpec(mj_copySpec(&other.model->spec), &other.model->spec.compiler); } // create a copy of the subtree that contains the frame @@ -2038,7 +2038,7 @@ mjCFrame& mjCFrame::operator=(const mjCFrame& other) { mjCFrame& mjCFrame::operator+=(const mjCBody& other) { // append a copy of the attached spec if (other.model != model && !model->FindSpec(mjs_getString(other.model->spec.modelname))) { - model->AppendSpec(mj_copySpec(&other.model->spec)); + model->AppendSpec(mj_copySpec(&other.model->spec), &other.model->spec.compiler); } // apply namespace and store keyframes in the source model diff --git a/test/user/user_api_test.cc b/test/user/user_api_test.cc index 84febf8e..3837c423 100644 --- a/test/user/user_api_test.cc +++ b/test/user/user_api_test.cc @@ -2514,6 +2514,54 @@ TEST_F(MujocoTest, DifferentUnitsAllowed) { mj_deleteModel(copied_model); } +TEST_F(MujocoTest, DifferentOptionsInAttachedFrame) { + static constexpr char xml_parent[] = R"( + + + + )"; + + static constexpr char xml_child[] = R"( + + + + + + + + + )"; + + // load specs and compile child + mjSpec* parent = mj_parseXMLString(xml_parent, 0, nullptr, 0); + EXPECT_THAT(parent, NotNull()); + mjSpec* child = mj_parseXMLString(xml_child, 0, nullptr, 0); + EXPECT_THAT(child, NotNull()); + mjModel* m_child = mj_compile(child, 0); + EXPECT_THAT(m_child, NotNull()); + + // attach child frame to parent worldbody + mjsBody* world = mjs_findBody(parent, "world"); + EXPECT_THAT(world, NotNull()); + mjsFrame* child_frame = mjs_findFrame(child, "child"); + EXPECT_THAT(child_frame, NotNull()); + mjsFrame* attached_frame = mjs_attachFrame(world, child_frame, "child-", ""); + EXPECT_THAT(attached_frame, NotNull()); + + // wrap the child frame in the parent frame and compile + mjModel* m_attached = mj_compile(parent, 0); + EXPECT_THAT(m_attached, NotNull()); + EXPECT_NEAR(m_attached->site_quat[0], m_child->site_quat[0], 1e-6); + EXPECT_NEAR(m_attached->site_quat[1], m_child->site_quat[1], 1e-6); + EXPECT_NEAR(m_attached->site_quat[2], m_child->site_quat[2], 1e-6); + EXPECT_NEAR(m_attached->site_quat[3], m_child->site_quat[3], 1e-6); + + mj_deleteSpec(parent); + mj_deleteSpec(child); + mj_deleteModel(m_child); + mj_deleteModel(m_attached); +} + TEST_F(MujocoTest, CopyAttachedSpec) { static constexpr char xml_parent[] = R"(