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
+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());