Add id and name properties to MjModel and MjData indexers.
PiperOrigin-RevId: 476078133 Change-Id: I905cc8ec147ac595cd56f0a42e38013ae1ac6903
This commit is contained in:
committed by
Copybara-Service
parent
6e256c03b7
commit
c13979cd17
@@ -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
@@ -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
@@ -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
|
||||
|
||||
@@ -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( \
|
||||
|
||||
Reference in New Issue
Block a user