Add mjs_findElement() to find any element in a spec by name.

Also, remove mjs_findMesh() and mjs_findKeyframe().

PiperOrigin-RevId: 670978980
Change-Id: Id88b800eb8a4c5efc40866f72a4442bf1d1437d0
This commit is contained in:
Alessio Quaglino
2024-09-04 08:23:18 -07:00
committed by Copybara-Service
parent 0e8c0b80eb
commit d3dfa6f970
13 changed files with 149 additions and 156 deletions
+10 -19
View File
@@ -4123,7 +4123,16 @@ mjs_findBody
.. mujoco-include:: mjs_findBody
Find body in model by name.
Find body in spec by name.
.. _mjs_findElement:
mjs_findElement
~~~~~~~~~~~~~~~
.. mujoco-include:: mjs_findElement
Find element in spec by name.
.. _mjs_findChild:
@@ -4134,15 +4143,6 @@ mjs_findChild
Find child body by name.
.. _mjs_findMesh:
mjs_findMesh
~~~~~~~~~~~~
.. mujoco-include:: mjs_findMesh
Find mesh by name.
.. _mjs_findFrame:
mjs_findFrame
@@ -4152,15 +4152,6 @@ mjs_findFrame
Find frame by name.
.. _mjs_findKeyframe:
mjs_findKeyframe
~~~~~~~~~~~~~~~~
.. mujoco-include:: mjs_findKeyframe
Find keyframe by name.
.. _mjs_getDefault:
mjs_getDefault
+2
View File
@@ -15,6 +15,8 @@ General
:ref:`ccd_tolerance<option-ccd_tolerance>` and :ref:`ccd_iterations<option-ccd_iterations>`, both in XML and in
the :ref:`mjOption` struct. This is because the new convex collision detection pipeline (see below) does not use
the MPR algorithm. The semantics of these options remain identical.
- The functions ``mjs_findMesh`` and ``mjs_findKeyframe`` were replaced by ``mjs_findElement``, which allows to look
for any object type.
- Added a new way of defining :ref:`connect<equality-connect>` equality constraints, using two sites rather than bodies.
The new semantic is useful when the assumption that the constraint is satisfied in the base configuration does not
+1 -2
View File
@@ -3572,10 +3572,9 @@ mjsTexture* mjs_addTexture(mjSpec* s);
mjsMaterial* mjs_addMaterial(mjSpec* s, mjsDefault* def);
mjSpec* mjs_getSpec(mjsBody* body);
mjsBody* mjs_findBody(mjSpec* s, const char* name);
mjsElement* mjs_findElement(mjSpec* s, mjtObj type, const char* name);
mjsBody* mjs_findChild(mjsBody* body, const char* name);
mjsMesh* mjs_findMesh(mjSpec* s, const char* name);
mjsFrame* mjs_findFrame(mjSpec* s, const char* name);
mjsKey* mjs_findKeyframe(mjSpec* s, const char* name);
mjsDefault* mjs_getDefault(mjsElement* element);
mjsDefault* mjs_findDefault(mjSpec* s, const char* classname);
mjsDefault* mjs_getSpecDefault(mjSpec* s);
+4 -7
View File
@@ -1522,21 +1522,18 @@ MJAPI mjsMaterial* mjs_addMaterial(mjSpec* s, mjsDefault* def);
// Get spec from body.
MJAPI mjSpec* mjs_getSpec(mjsBody* body);
// Find body in model by name.
// Find body in spec by name.
MJAPI mjsBody* mjs_findBody(mjSpec* s, const char* name);
// Find element in spec by name.
MJAPI mjsElement* mjs_findElement(mjSpec* s, mjtObj type, const char* name);
// Find child body by name.
MJAPI mjsBody* mjs_findChild(mjsBody* body, const char* name);
// Find mesh by name.
MJAPI mjsMesh* mjs_findMesh(mjSpec* s, const char* name);
// Find frame by name.
MJAPI mjsFrame* mjs_findFrame(mjSpec* s, const char* name);
// Find keyframe by name.
MJAPI mjsKey* mjs_findKeyframe(mjSpec* s, const char* name);
// Get default corresponding to an element.
MJAPI mjsDefault* mjs_getDefault(mjsElement* element);
+27 -45
View File
@@ -9678,7 +9678,33 @@ FUNCTIONS: Mapping[str, FunctionDecl] = dict([
),
),
),
doc='Find body in model by name.',
doc='Find body in spec by name.',
)),
('mjs_findElement',
FunctionDecl(
name='mjs_findElement',
return_type=PointerType(
inner_type=ValueType(name='mjsElement'),
),
parameters=(
FunctionParameterDecl(
name='s',
type=PointerType(
inner_type=ValueType(name='mjSpec'),
),
),
FunctionParameterDecl(
name='type',
type=ValueType(name='mjtObj'),
),
FunctionParameterDecl(
name='name',
type=PointerType(
inner_type=ValueType(name='char', is_const=True),
),
),
),
doc='Find element in spec by name.',
)),
('mjs_findChild',
FunctionDecl(
@@ -9702,28 +9728,6 @@ FUNCTIONS: Mapping[str, FunctionDecl] = dict([
),
doc='Find child body by name.',
)),
('mjs_findMesh',
FunctionDecl(
name='mjs_findMesh',
return_type=PointerType(
inner_type=ValueType(name='mjsMesh'),
),
parameters=(
FunctionParameterDecl(
name='s',
type=PointerType(
inner_type=ValueType(name='mjSpec'),
),
),
FunctionParameterDecl(
name='name',
type=PointerType(
inner_type=ValueType(name='char', is_const=True),
),
),
),
doc='Find mesh by name.',
)),
('mjs_findFrame',
FunctionDecl(
name='mjs_findFrame',
@@ -9746,28 +9750,6 @@ FUNCTIONS: Mapping[str, FunctionDecl] = dict([
),
doc='Find frame by name.',
)),
('mjs_findKeyframe',
FunctionDecl(
name='mjs_findKeyframe',
return_type=PointerType(
inner_type=ValueType(name='mjsKey'),
),
parameters=(
FunctionParameterDecl(
name='s',
type=PointerType(
inner_type=ValueType(name='mjSpec'),
),
),
FunctionParameterDecl(
name='name',
type=PointerType(
inner_type=ValueType(name='char', is_const=True),
),
),
),
doc='Find keyframe by name.',
)),
('mjs_getDefault',
FunctionDecl(
name='mjs_getDefault',
-12
View File
@@ -133,24 +133,12 @@ PYBIND11_MODULE(_specs, m) {
return mjs_findBody(self.ptr, name.c_str());
},
py::return_value_policy::reference_internal);
mjSpec.def(
"find_mesh",
[](MjSpec& self, std::string& name) -> raw::MjsMesh* {
return mjs_findMesh(self.ptr, name.c_str());
},
py::return_value_policy::reference_internal);
mjSpec.def(
"find_frame",
[](MjSpec& self, std::string& name) -> raw::MjsFrame* {
return mjs_findFrame(self.ptr, name.c_str());
},
py::return_value_policy::reference_internal);
mjSpec.def(
"find_keyframe",
[](MjSpec& self, std::string& name) -> raw::MjsKey* {
return mjs_findKeyframe(self.ptr, name.c_str());
},
py::return_value_policy::reference_internal);
mjSpec.def(
"find_default",
[](MjSpec& self, std::string& classname) -> raw::MjsDefault* {
+24 -27
View File
@@ -542,18 +542,34 @@ mjsDefault* mjs_getSpecDefault(mjSpec* s) {
// find body in model by name
mjsBody* mjs_findBody(mjSpec* s, const char* name) {
mjCModel* model = static_cast<mjCModel*>(s->element);
mjCBase* body = 0;
if (model->IsCompiled()) {
body = model->FindObject(mjOBJ_BODY, std::string(name)); // fast lookup
} else {
body = model->FindBody(model->GetWorld(), std::string(name)); // recursive search
}
mjsElement* body = mjs_findElement(s, mjOBJ_BODY, name);
return body ? &(static_cast<mjCBody*>(body)->spec) : nullptr;
}
// find element in spec by name
mjsElement* mjs_findElement(mjSpec* s, mjtObj type, const char* name) {
mjCModel* model = static_cast<mjCModel*>(s->element);
if (model->IsCompiled()) {
return model->FindObject(type, std::string(name)); // fast lookup
}
switch (type) {
case mjOBJ_BODY:
case mjOBJ_SITE:
case mjOBJ_GEOM:
case mjOBJ_JOINT:
case mjOBJ_CAMERA:
case mjOBJ_LIGHT:
case mjOBJ_FRAME:
return model->FindTree(model->GetWorld(), type, std::string(name)); // recursive search
default:
return model->FindObject(type, std::string(name)); // always available
}
}
// find child of a body by name
mjsBody* mjs_findChild(mjsBody* bodyspec, const char* name) {
mjCBody* body = static_cast<mjCBody*>(bodyspec->element);
@@ -563,33 +579,14 @@ mjsBody* mjs_findChild(mjsBody* bodyspec, const char* name) {
// find mesh by name
mjsMesh* mjs_findMesh(mjSpec* s, const char* name) {
mjCModel* model = static_cast<mjCModel*>(s->element);
mjCMesh* mesh = (mjCMesh*)model->FindObject(mjOBJ_MESH, std::string(name));
return mesh ? &(static_cast<mjCMesh*>(mesh)->spec) : nullptr;
}
// find frame by name
mjsFrame* mjs_findFrame(mjSpec* s, const char* name) {
mjCModel* model = static_cast<mjCModel*>(s->element);
mjCFrame* frame = (mjCFrame*)model->FindFrame(model->GetWorld(), std::string(name));
mjsElement* frame = mjs_findElement(s, mjOBJ_FRAME, name);
return frame ? &(static_cast<mjCFrame*>(frame)->spec) : nullptr;
}
// find keyframe by name
mjsKey* mjs_findKeyframe(mjSpec* s, const char* name) {
mjCModel* model = static_cast<mjCModel*>(s->element);
mjCKey* key = (mjCKey*)model->FindObject(mjOBJ_KEY, std::string(name));
return key ? &(static_cast<mjCKey*>(key)->spec) : nullptr;
}
// set frame
void mjs_setFrame(mjsElement* dest, mjsFrame* frame) {
if (!frame) {
+4 -7
View File
@@ -188,21 +188,18 @@ MJAPI mjSpec* mjs_getSpec(mjsBody* body);
// Find spec (model asset) by name.
MJAPI mjSpec* mjs_findSpec(mjSpec* spec, const char* name);
// Find body in model by name.
// Find body in spec by name.
MJAPI mjsBody* mjs_findBody(mjSpec* s, const char* name);
// Find element in spec by name.
MJAPI mjsElement* mjs_findElement(mjSpec* s, mjtObj type, const char* name);
// Find child body by name.
MJAPI mjsBody* mjs_findChild(mjsBody* body, const char* name);
// Find mesh by name.
MJAPI mjsMesh* mjs_findMesh(mjSpec* s, const char* name);
// Find frame by name.
MJAPI mjsFrame* mjs_findFrame(mjSpec* s, const char* name);
// Find keyframe by name.
MJAPI mjsKey* mjs_findKeyframe(mjSpec* s, const char* name);
// Get default corresponding to an element.
MJAPI mjsDefault* mjs_getDefault(mjsElement* element);
+52 -24
View File
@@ -1031,33 +1031,61 @@ mjCBase* mjCModel::FindObject(mjtObj type, string name) const {
// find body by name
mjCBody* mjCModel::FindBody(mjCBody* body, std::string name) {
if (body->name == name) {
return body;
mjCBase* mjCModel::FindTree(mjCBody* body, mjtObj type, std::string name) {
switch (type) {
case mjOBJ_BODY:
if (body->name == name) {
return body;
}
break;
case mjOBJ_SITE:
for (auto site : body->sites) {
if (site->name == name) {
return site;
}
}
break;
case mjOBJ_GEOM:
for (auto geom : body->geoms) {
if (geom->name == name) {
return geom;
}
}
break;
case mjOBJ_JOINT:
for (auto joint : body->joints) {
if (joint->name == name) {
return joint;
}
}
break;
case mjOBJ_CAMERA:
for (auto camera : body->cameras) {
if (camera->name == name) {
return camera;
}
}
break;
case mjOBJ_LIGHT:
for (auto light : body->lights) {
if (light->name == name) {
return light;
}
}
break;
case mjOBJ_FRAME:
for (auto frame : body->frames) {
if (frame->name == name) {
return frame;
}
}
break;
default:
return nullptr;
}
for (auto child : body->bodies) {
auto candidate = FindBody(child, name);
if (candidate) {
return candidate;
}
}
return nullptr;
}
// find frame by name
mjCFrame* mjCModel::FindFrame(mjCBody* body, std::string name) const{
for (auto frame : body->frames) {
if (frame->name == name) {
return frame;
}
}
for (auto body : body->bodies) {
auto candidate = FindFrame(body, name);
auto candidate = FindTree(child, type, name);
if (candidate) {
return candidate;
}
+9 -10
View File
@@ -223,16 +223,15 @@ class mjCModel : public mjCModel_, private mjSpec {
mjsElement* NextObject(mjsElement* object, mjtObj type = mjOBJ_UNKNOWN); // next object of specified type
// API for access to other variables
bool IsCompiled() const; // is model already compiled
const mjCError& GetError() const; // get reference of error object
void SetError(const mjCError& error) { errInfo = error; } // set value of error object
mjCBody* GetWorld(); // pointer to world body
mjCDef* FindDefault(std::string name); // find defaults class name
mjCDef* AddDefault(std::string name, mjCDef* parent = nullptr); // add defaults class to array
mjCBase* FindObject(mjtObj type, std::string name) const; // find object given type and name
mjCBody* FindBody(mjCBody* body, std::string name); // find body given name
mjCFrame* FindFrame(mjCBody* body, std::string name) const; // find frame given name
mjSpec* FindSpec(std::string name) const; // find spec given name
bool IsCompiled() const; // is model already compiled
const mjCError& GetError() const; // get reference of error object
void SetError(const mjCError& error) { errInfo = error; } // set value of error object
mjCBody* GetWorld(); // pointer to world body
mjCDef* FindDefault(std::string name); // find defaults class name
mjCDef* AddDefault(std::string name, mjCDef* parent = nullptr); // add defaults class to array
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
void SetActivePlugins(const std::vector<std::pair<const mjpPlugin*, int>>&& active_plugins) {
active_plugins_ = std::move(active_plugins);
}
+2 -2
View File
@@ -1408,8 +1408,8 @@ void mjCBody::ForgetKeyframes() const {
joint->qpos_.clear();
joint->qvel_.clear();
}
model->FindBody((mjCBody*)this, name)->mpos_.clear(); // this is a hack to avoid const
model->FindBody((mjCBody*)this, name)->mquat_.clear(); // this is a hack to avoid const
((mjCBody*)this)->mpos_.clear();
((mjCBody*)this)->mquat_.clear();
for (auto body : bodies) {
body->ForgetKeyframes();
}
+1 -1
View File
@@ -592,7 +592,7 @@ mjsGeom* mjXURDF::Geom(XMLElement* geom_elem, mjsBody* pbody, bool collision) {
meshname = mjuu_stripext(meshname);
// look for existing mesh
mjsMesh* mesh = mjs_findMesh(spec, meshname.c_str());
mjsMesh* mesh = mjs_asMesh(mjs_findElement(spec, mjOBJ_MESH, meshname.c_str()));
mjsMesh* pmesh = 0;
// does not exist: create
+13
View File
@@ -78,6 +78,13 @@ TEST_F(MujocoTest, TreeTraversal) {
mjsSite* site3 = mjs_addSite(body, 0);
mjsGeom* geom3 = mjs_addGeom(body, 0);
mjs_setString(site1->name, "site1");
mjs_setString(geom1->name, "geom1");
mjs_setString(geom2->name, "geom2");
mjs_setString(site2->name, "site2");
mjs_setString(site3->name, "site3");
mjs_setString(geom3->name, "geom3");
mjsElement* a_el1 = mjs_firstElement(spec, mjOBJ_ACTUATOR);
mjsElement* c_el1 = mjs_firstChild(body, mjOBJ_CAMERA);
mjsElement* t_el1 = mjs_firstChild(body, mjOBJ_TENDON);
@@ -101,6 +108,12 @@ TEST_F(MujocoTest, TreeTraversal) {
EXPECT_EQ(g_el3, geom3->element);
EXPECT_EQ(g_el4, nullptr);
EXPECT_EQ(s_el4, nullptr);
EXPECT_EQ(mjs_findElement(spec, mjOBJ_SITE, "site1"), site1->element);
EXPECT_EQ(mjs_findElement(spec, mjOBJ_SITE, "site2"), site2->element);
EXPECT_EQ(mjs_findElement(spec, mjOBJ_SITE, "site3"), site3->element);
EXPECT_EQ(mjs_findElement(spec, mjOBJ_GEOM, "geom1"), geom1->element);
EXPECT_EQ(mjs_findElement(spec, mjOBJ_GEOM, "geom2"), geom2->element);
EXPECT_EQ(mjs_findElement(spec, mjOBJ_GEOM, "geom3"), geom3->element);
mj_deleteSpec(spec);
}