Add id and name properties to MjModel and MjData indexers.

PiperOrigin-RevId: 476078133
Change-Id: I905cc8ec147ac595cd56f0a42e38013ae1ac6903
This commit is contained in:
Saran Tunyasuvunakool
2022-09-22 06:18:20 -07:00
committed by Copybara-Service
parent 6e256c03b7
commit c13979cd17
6 changed files with 213 additions and 61 deletions
+44
View File
@@ -265,6 +265,8 @@ class MuJoCoBindingsTest(parameterized.TestCase):
def test_named_indexing_repr_in_data(self):
expected_repr = '''<_MjDataGeomViews
id: 1
name: 'mybox'
xmat: array([0., 0., 0., 0., 0., 0., 0., 0., 0.])
xpos: array([0., 0., 0.])
>'''
@@ -1038,6 +1040,48 @@ Euler integrator, semi-implicit in velocity.
)
self._assert_attributes_equal(model2, self.model, attr_to_compare)
def test_indexer_name_id(self):
xml = r"""
<mujoco>
<worldbody>
<geom name="mygeom" size="1" pos="0 0 1"/>
<geom size="2" pos="0 0 2"/>
<geom size="3" pos="0 0 3"/>
<geom name="myothergeom" size="4" pos="0 0 4"/>
<geom size="5" pos="0 0 5"/>
</worldbody>
</mujoco>
"""
model = mujoco.MjModel.from_xml_string(xml)
self.assertEqual(model.geom('mygeom').id, 0)
self.assertEqual(model.geom('myothergeom').id, 3)
self.assertEqual(model.geom(0).name, 'mygeom')
self.assertEqual(model.geom(1).name, '')
self.assertEqual(model.geom(2).name, '')
self.assertEqual(model.geom(3).name, 'myothergeom')
self.assertEqual(model.geom(4).name, '')
self.assertEqual(model.geom(0).size[0], 1)
self.assertEqual(model.geom(1).size[0], 2)
self.assertEqual(model.geom(2).size[0], 3)
self.assertEqual(model.geom(3).size[0], 4)
self.assertEqual(model.geom(4).size[0], 5)
data = mujoco.MjData(model)
mujoco.mj_forward(model, data)
self.assertEqual(data.geom('mygeom').id, 0)
self.assertEqual(data.geom('myothergeom').id, 3)
self.assertEqual(data.geom(0).name, 'mygeom')
self.assertEqual(data.geom(1).name, '')
self.assertEqual(data.geom(2).name, '')
self.assertEqual(data.geom(3).name, 'myothergeom')
self.assertEqual(data.geom(4).name, '')
self.assertEqual(data.geom(0).xpos[2], 1)
self.assertEqual(data.geom(1).xpos[2], 2)
self.assertEqual(data.geom(2).xpos[2], 3)
self.assertEqual(data.geom(3).xpos[2], 4)
self.assertEqual(data.geom(4).xpos[2], 5)
def _assert_attributes_equal(self, actual_obj, expected_obj, attr_to_compare):
for name in attr_to_compare:
actual_value = getattr(actual_obj, name)
+75 -29
View File
@@ -39,9 +39,9 @@ namespace {
// each name in the `names` array.
// names: Character arrays consisting of the concatenation of all names.
template <typename IntPtr, typename CharPtr>
NameToIDMap MakeMap(int count, IntPtr name_offsets, CharPtr names)
NameToID MakeNameToID(int count, IntPtr name_offsets, CharPtr names)
{
NameToIDMap name_to_id;
NameToID name_to_id;
for (int index = 0; index < count; ++index) {
const char* name = &names[name_offsets[index]];
if (name[0] != '\0') {
@@ -51,6 +51,23 @@ NameToIDMap MakeMap(int count, IntPtr name_offsets, CharPtr names)
return name_to_id;
}
// Parses raw mjModel to create a mapping from array indices to names.
//
// Args:
// count: Number of names in the map.
// name_offsets: Array consisting of indices that correspond to the start of
// each name in the `names` array.
// names: Character arrays consisting of the concatenation of all names.
template <typename IntPtr, typename CharPtr>
IDToName MakeIDToName(int count, IntPtr name_offsets, CharPtr names) {
IDToName id_to_name;
for (int index = 0; index < count; ++index) {
const char* name = &names[name_offsets[index]];
id_to_name.emplace_back(name);
}
return id_to_name;
};
// Makes an array view into an mjModel/mjData struct field at a given index.
//
// Template args:
@@ -120,33 +137,59 @@ py::array_t<T> MakeArray(T* base_ptr, int index, std::vector<int>&& shape,
// M is either a raw::MjModel or MjDataMetadata.
template <typename M>
NameToIDMaps::NameToIDMaps(const M& m)
: body(MakeMap(m.nbody, m.name_bodyadr, m.names)),
jnt(MakeMap(m.njnt, m.name_jntadr, m.names)),
geom(MakeMap(m.ngeom, m.name_geomadr, m.names)),
site(MakeMap(m.nsite, m.name_siteadr, m.names)),
cam(MakeMap(m.ncam, m.name_camadr, m.names)),
light(MakeMap(m.nlight, m.name_lightadr, m.names)),
mesh(MakeMap(m.nmesh, m.name_meshadr, m.names)),
skin(MakeMap(m.nskin, m.name_skinadr, m.names)),
hfield(MakeMap(m.nhfield, m.name_hfieldadr, m.names)),
tex(MakeMap(m.ntex, m.name_texadr, m.names)),
mat(MakeMap(m.nmat, m.name_matadr, m.names)),
pair(MakeMap(m.npair, m.name_pairadr, m.names)),
exclude(MakeMap(m.nexclude, m.name_excludeadr, m.names)),
eq(MakeMap(m.neq, m.name_eqadr, m.names)),
tendon(MakeMap(m.ntendon, m.name_tendonadr, m.names)),
actuator(MakeMap(m.nu, m.name_actuatoradr, m.names)),
sensor(MakeMap(m.nsensor, m.name_sensoradr, m.names)),
numeric(MakeMap(m.nnumeric, m.name_numericadr, m.names)),
text(MakeMap(m.ntext, m.name_textadr, m.names)),
tuple(MakeMap(m.ntuple, m.name_tupleadr, m.names)),
key(MakeMap(m.nkey, m.name_keyadr, m.names)) {}
NameToIDMappings::NameToIDMappings(const M& m)
: body(MakeNameToID(m.nbody, m.name_bodyadr, m.names)),
jnt(MakeNameToID(m.njnt, m.name_jntadr, m.names)),
geom(MakeNameToID(m.ngeom, m.name_geomadr, m.names)),
site(MakeNameToID(m.nsite, m.name_siteadr, m.names)),
cam(MakeNameToID(m.ncam, m.name_camadr, m.names)),
light(MakeNameToID(m.nlight, m.name_lightadr, m.names)),
mesh(MakeNameToID(m.nmesh, m.name_meshadr, m.names)),
skin(MakeNameToID(m.nskin, m.name_skinadr, m.names)),
hfield(MakeNameToID(m.nhfield, m.name_hfieldadr, m.names)),
tex(MakeNameToID(m.ntex, m.name_texadr, m.names)),
mat(MakeNameToID(m.nmat, m.name_matadr, m.names)),
pair(MakeNameToID(m.npair, m.name_pairadr, m.names)),
exclude(MakeNameToID(m.nexclude, m.name_excludeadr, m.names)),
eq(MakeNameToID(m.neq, m.name_eqadr, m.names)),
tendon(MakeNameToID(m.ntendon, m.name_tendonadr, m.names)),
actuator(MakeNameToID(m.nu, m.name_actuatoradr, m.names)),
sensor(MakeNameToID(m.nsensor, m.name_sensoradr, m.names)),
numeric(MakeNameToID(m.nnumeric, m.name_numericadr, m.names)),
text(MakeNameToID(m.ntext, m.name_textadr, m.names)),
tuple(MakeNameToID(m.ntuple, m.name_tupleadr, m.names)),
key(MakeNameToID(m.nkey, m.name_keyadr, m.names)) {}
// M is either a raw::MjModel or MjDataMetadata.
template <typename M>
IDToNameMappings::IDToNameMappings(const M& m)
: body(MakeIDToName(m.nbody, m.name_bodyadr, m.names)),
jnt(MakeIDToName(m.njnt, m.name_jntadr, m.names)),
geom(MakeIDToName(m.ngeom, m.name_geomadr, m.names)),
site(MakeIDToName(m.nsite, m.name_siteadr, m.names)),
cam(MakeIDToName(m.ncam, m.name_camadr, m.names)),
light(MakeIDToName(m.nlight, m.name_lightadr, m.names)),
mesh(MakeIDToName(m.nmesh, m.name_meshadr, m.names)),
skin(MakeIDToName(m.nskin, m.name_skinadr, m.names)),
hfield(MakeIDToName(m.nhfield, m.name_hfieldadr, m.names)),
tex(MakeIDToName(m.ntex, m.name_texadr, m.names)),
mat(MakeIDToName(m.nmat, m.name_matadr, m.names)),
pair(MakeIDToName(m.npair, m.name_pairadr, m.names)),
exclude(MakeIDToName(m.nexclude, m.name_excludeadr, m.names)),
eq(MakeIDToName(m.neq, m.name_eqadr, m.names)),
tendon(MakeIDToName(m.ntendon, m.name_tendonadr, m.names)),
actuator(MakeIDToName(m.nu, m.name_actuatoradr, m.names)),
sensor(MakeIDToName(m.nsensor, m.name_sensoradr, m.names)),
numeric(MakeIDToName(m.nnumeric, m.name_numericadr, m.names)),
text(MakeIDToName(m.ntext, m.name_textadr, m.names)),
tuple(MakeIDToName(m.ntuple, m.name_tupleadr, m.names)),
key(MakeIDToName(m.nkey, m.name_keyadr, m.names)) {}
MjModelIndexer::MjModelIndexer(raw::MjModel* m, py::handle owner)
: m_(m),
owner_(owner),
name_to_id_(*m)
name_to_id_(*m),
id_to_name_(*m)
#define XGROUP(MjModelFieldGroupedViews, field, nfield, FIELD_XMACROS) \
, field##_(m->nfield, std::nullopt)
MJMODEL_VIEW_GROUPS
@@ -160,7 +203,8 @@ MjModelIndexer::MjModelIndexer(raw::MjModel* m, py::handle owner)
} \
auto& indexer = field##_[i]; \
if (!indexer.has_value()) { \
indexer.emplace(i, m_, owner_); \
const std::string& name = id_to_name_.field[i]; \
indexer.emplace(i, name, m_, owner_); \
} \
return *indexer; \
}
@@ -184,7 +228,8 @@ MjDataIndexer::MjDataIndexer(raw::MjData* d, const MjDataMetadata* m,
: d_(d),
m_(m),
owner_(owner),
name_to_id_(*m)
name_to_id_(*m),
id_to_name_(*m)
#define XGROUP(MjDataGroupedViews, field, nfield, FIELD_XMACROS) \
, field##_(m->nfield, std::nullopt)
MJDATA_VIEW_GROUPS
@@ -198,7 +243,8 @@ MjDataIndexer::MjDataIndexer(raw::MjData* d, const MjDataMetadata* m,
} \
auto& indexer = field##_[i]; \
if (!indexer.has_value()) { \
indexer.emplace(i, d_, m_, owner_); \
const std::string& name = id_to_name_.field[i]; \
indexer.emplace(i, name, d_, m_, owner_); \
} \
return *indexer; \
}
@@ -371,7 +417,7 @@ MJDATA_TENDON
// Returns an error message when a non-existent name is requested from an
// indexer, which includes all valid names.
std::string KeyErrorMessage(const NameToIDMap& map, std::string_view name) {
std::string KeyErrorMessage(const NameToID& map, std::string_view name) {
// Make a sorted list of valid names
std::vector<std::string_view> valid_names;
valid_names.reserve(map.size());
+70 -31
View File
@@ -31,34 +31,63 @@
#include <pybind11/pybind11.h>
namespace mujoco::python {
using NameToIDMap = absl::flat_hash_map<std::string, int>;
using NameToID = absl::flat_hash_map<std::string, int>;
using IDToName = std::vector<std::string>;
struct NameToIDMaps {
struct NameToIDMappings {
// M is either a raw::MjModel or MjDataMetadata.
template <typename M>
explicit NameToIDMaps(const M& m);
explicit NameToIDMappings(const M& m);
NameToIDMap body;
NameToIDMap jnt;
NameToIDMap geom;
NameToIDMap site;
NameToIDMap cam;
NameToIDMap light;
NameToIDMap mesh;
NameToIDMap skin;
NameToIDMap hfield;
NameToIDMap tex;
NameToIDMap mat;
NameToIDMap pair;
NameToIDMap exclude;
NameToIDMap eq;
NameToIDMap tendon;
NameToIDMap actuator;
NameToIDMap sensor;
NameToIDMap numeric;
NameToIDMap text;
NameToIDMap tuple;
NameToIDMap key;
NameToID body;
NameToID jnt;
NameToID geom;
NameToID site;
NameToID cam;
NameToID light;
NameToID mesh;
NameToID skin;
NameToID hfield;
NameToID tex;
NameToID mat;
NameToID pair;
NameToID exclude;
NameToID eq;
NameToID tendon;
NameToID actuator;
NameToID sensor;
NameToID numeric;
NameToID text;
NameToID tuple;
NameToID key;
};
struct IDToNameMappings {
// M is either a raw::MjModel or MjDataMetadata.
template <typename M>
explicit IDToNameMappings(const M& m);
IDToName body;
IDToName jnt;
IDToName geom;
IDToName site;
IDToName cam;
IDToName light;
IDToName mesh;
IDToName skin;
IDToName hfield;
IDToName tex;
IDToName mat;
IDToName pair;
IDToName exclude;
IDToName eq;
IDToName tendon;
IDToName actuator;
IDToName sensor;
IDToName numeric;
IDToName text;
IDToName tuple;
IDToName key;
};
@@ -67,12 +96,16 @@ struct NameToIDMaps {
// geom, or a particular joint).
class MjModelGroupedViewsBase {
public:
MjModelGroupedViewsBase(int index, raw::MjModel* m,
MjModelGroupedViewsBase(int index, std::string_view name, raw::MjModel* m,
pybind11::handle owner)
: index_(index), m_(m), owner_(owner) {}
: index_(index), name_(name), m_(m), owner_(owner) {}
int index() const { return index_; }
const std::string& name() const { return name_; }
protected:
int index_;
std::string name_;
raw::MjModel* m_;
pybind11::handle owner_;
};
@@ -111,7 +144,8 @@ class MjModelIndexer {
private:
raw::MjModel* m_;
pybind11::handle owner_;
NameToIDMaps name_to_id_;
NameToIDMappings name_to_id_;
IDToNameMappings id_to_name_;
// Lazily instantiate a grouped views object when accessed from Python, but
// cache it once made so that we can return the same one if requested again.
@@ -127,13 +161,17 @@ class MjModelIndexer {
// geom, or a particular joint).
class MjDataGroupedViewsBase {
public:
MjDataGroupedViewsBase(int index, raw::MjData* d,
MjDataGroupedViewsBase(int index, std::string_view name, raw::MjData* d,
const MjDataMetadata* m,
pybind11::handle owner)
: index_(index), d_(d), m_(m), owner_(owner) {}
: index_(index), name_(name), d_(d), m_(m), owner_(owner) {}
int index() const { return index_; }
const std::string& name() const { return name_; }
protected:
int index_;
std::string name_;
raw::MjData* d_;
const MjDataMetadata* m_;
pybind11::handle owner_;
@@ -175,7 +213,8 @@ class MjDataIndexer {
raw::MjData* d_;
const MjDataMetadata* m_;
pybind11::handle owner_;
NameToIDMaps name_to_id_;
NameToIDMappings name_to_id_;
IDToNameMappings id_to_name_;
// Lazily instantiate a grouped views object when accessed from Python, but
// cache it once made so that we can return the same one if requested again.
@@ -188,7 +227,7 @@ class MjDataIndexer {
// Returns an error message when a nonexistent name is requested from an
// indexer, which includes all valid names.
std::string KeyErrorMessage(const NameToIDMap& map, std::string_view name);
std::string KeyErrorMessage(const NameToID& map, std::string_view name);
// Returns an error message when an invalid numeric index is provided.
std::string IndexErrorMessage(int index, int size);
} // namespace mujoco::python
+12
View File
@@ -1506,6 +1506,12 @@ This is useful for example when the MJB is not available as a file on disk.)"));
py::class_<MjModelGroupedViews> groupedViews(m, "_" #MjModelGroupedViews); \
FIELD_XMACROS \
groupedViews.def("__repr__", StructRepr<GroupedViews>); \
groupedViews.def_property_readonly("id", [](GroupedViews& views) { \
return views.index(); \
}); \
groupedViews.def_property_readonly("name", [](GroupedViews& views) { \
return views.name(); \
}); \
}
#define X(type, prefix, var, dim0, dim1) \
groupedViews.def_property( \
@@ -1780,6 +1786,12 @@ This is useful for example when the MJB is not available as a file on disk.)"));
py::class_<MjDataGroupedViews> groupedViews(m, "_" #MjDataGroupedViews); \
FIELD_XMACROS \
groupedViews.def("__repr__", StructRepr<GroupedViews>); \
groupedViews.def_property_readonly("id", [](GroupedViews& views) { \
return views.index(); \
}); \
groupedViews.def_property_readonly("name", [](GroupedViews& views) { \
return views.name(); \
}); \
}
#define X(type, prefix, var, dim0, dim1) \
groupedViews.def_property( \