Modify mjListKeyMap to be a fixed-size array, and add internal object_lists variable to mjCModel to remove redundant code.

PiperOrigin-RevId: 590997348
Change-Id: I5f4d7de605f44d40a75eb1e9ceec208325fa3be5
This commit is contained in:
Kyle Bayes
2023-12-14 11:11:00 -08:00
committed by Copybara-Service
parent 0915d69c3f
commit 9f6be3135e
3 changed files with 80 additions and 187 deletions
+50 -184
View File
@@ -147,7 +147,6 @@ mjCModel::mjCModel() {
nuser_sensor = -1;
//------------------------ private variables
ids.clear();
cameras.clear();
lights.clear();
flexes.clear();
@@ -183,6 +182,35 @@ mjCModel::mjCModel() {
world->name = "world";
world->def = defaults[0];
bodies.push_back(world);
for (int i = 0; i < mjNOBJECT; ++i) {
object_lists[i] = nullptr;
}
object_lists[mjOBJ_BODY] = (std::vector<mjCBase*>*) &bodies;
object_lists[mjOBJ_XBODY] = (std::vector<mjCBase*>*) &bodies;
object_lists[mjOBJ_JOINT] = (std::vector<mjCBase*>*) &joints;
object_lists[mjOBJ_GEOM] = (std::vector<mjCBase*>*) &geoms;
object_lists[mjOBJ_SITE] = (std::vector<mjCBase*>*) &sites;
object_lists[mjOBJ_CAMERA] = (std::vector<mjCBase*>*) &cameras;
object_lists[mjOBJ_LIGHT] = (std::vector<mjCBase*>*) &lights;
object_lists[mjOBJ_FLEX] = (std::vector<mjCBase*>*) &flexes;
object_lists[mjOBJ_MESH] = (std::vector<mjCBase*>*) &meshes;
object_lists[mjOBJ_SKIN] = (std::vector<mjCBase*>*) &skins;
object_lists[mjOBJ_HFIELD] = (std::vector<mjCBase*>*) &hfields;
object_lists[mjOBJ_TEXTURE] = (std::vector<mjCBase*>*) &textures;
object_lists[mjOBJ_MATERIAL] = (std::vector<mjCBase*>*) &materials;
object_lists[mjOBJ_PAIR] = (std::vector<mjCBase*>*) &pairs;
object_lists[mjOBJ_EXCLUDE] = (std::vector<mjCBase*>*) &excludes;
object_lists[mjOBJ_EQUALITY] = (std::vector<mjCBase*>*) &equalities;
object_lists[mjOBJ_TENDON] = (std::vector<mjCBase*>*) &tendons;
object_lists[mjOBJ_ACTUATOR] = (std::vector<mjCBase*>*) &actuators;
object_lists[mjOBJ_SENSOR] = (std::vector<mjCBase*>*) &sensors;
object_lists[mjOBJ_NUMERIC] = (std::vector<mjCBase*>*) &numerics;
object_lists[mjOBJ_TEXT] = (std::vector<mjCBase*>*) &texts;
object_lists[mjOBJ_TUPLE] = (std::vector<mjCBase*>*) &tuples;
object_lists[mjOBJ_KEY] = (std::vector<mjCBase*>*) &keys;
object_lists[mjOBJ_PLUGIN] = (std::vector<mjCBase*>*) &plugins;
}
@@ -213,7 +241,6 @@ mjCModel::~mjCModel() {
for (int i=0; i<defaults.size(); i++) delete defaults[i];
// clear pointer lists created in model construction
ids.clear();
flexes.clear();
meshes.clear();
skins.clear();
@@ -456,118 +483,20 @@ mjCPlugin* mjCModel::AddPlugin(void) {
// get number of objects of specified type
int mjCModel::NumObjects(mjtObj type) {
switch (type) {
case mjOBJ_BODY:
case mjOBJ_XBODY:
return (int)bodies.size();
case mjOBJ_JOINT:
return (int)joints.size();
case mjOBJ_GEOM:
return (int)geoms.size();
case mjOBJ_SITE:
return (int)sites.size();
case mjOBJ_CAMERA:
return (int)cameras.size();
case mjOBJ_LIGHT:
return (int)lights.size();
case mjOBJ_FLEX:
return (int)flexes.size();
case mjOBJ_MESH:
return (int)meshes.size();
case mjOBJ_SKIN:
return (int)skins.size();
case mjOBJ_HFIELD:
return (int)hfields.size();
case mjOBJ_TEXTURE:
return (int)textures.size();
case mjOBJ_MATERIAL:
return (int)materials.size();
case mjOBJ_PAIR:
return (int)pairs.size();
case mjOBJ_EXCLUDE:
return (int)excludes.size();
case mjOBJ_EQUALITY:
return (int)equalities.size();
case mjOBJ_TENDON:
return (int)tendons.size();
case mjOBJ_ACTUATOR:
return (int)actuators.size();
case mjOBJ_SENSOR:
return (int)sensors.size();
case mjOBJ_NUMERIC:
return (int)numerics.size();
case mjOBJ_TEXT:
return (int)texts.size();
case mjOBJ_TUPLE:
return (int)tuples.size();
case mjOBJ_KEY:
return (int)keys.size();
case mjOBJ_PLUGIN:
return (int)plugins.size();
default:
if (!object_lists[type]) {
return 0;
}
return (int) object_lists[type]->size();
}
// get pointer to specified object
mjCBase* mjCModel::GetObject(mjtObj type, int id) {
if (id>=0 && id<NumObjects(type)) {
switch (type) {
case mjOBJ_BODY:
case mjOBJ_XBODY:
return bodies[id];
case mjOBJ_JOINT:
return joints[id];
case mjOBJ_GEOM:
return geoms[id];
case mjOBJ_SITE:
return sites[id];
case mjOBJ_CAMERA:
return cameras[id];
case mjOBJ_LIGHT:
return lights[id];
case mjOBJ_FLEX:
return flexes[id];
case mjOBJ_MESH:
return meshes[id];
case mjOBJ_SKIN:
return skins[id];
case mjOBJ_HFIELD:
return hfields[id];
case mjOBJ_TEXTURE:
return textures[id];
case mjOBJ_MATERIAL:
return materials[id];
case mjOBJ_PAIR:
return pairs[id];
case mjOBJ_EXCLUDE:
return excludes[id];
case mjOBJ_EQUALITY:
return equalities[id];
case mjOBJ_TENDON:
return tendons[id];
case mjOBJ_ACTUATOR:
return actuators[id];
case mjOBJ_SENSOR:
return sensors[id];
case mjOBJ_NUMERIC:
return numerics[id];
case mjOBJ_TEXT:
return texts[id];
case mjOBJ_TUPLE:
return tuples[id];
case mjOBJ_KEY:
return keys[id];
case mjOBJ_PLUGIN:
return plugins[id];
default:
return 0;
}
if (id < 0 || id >= NumObjects(type)) {
return nullptr;
}
return 0;
return (*object_lists[type])[id];
}
@@ -666,55 +595,10 @@ static T* findobject(std::string_view name, const vector<T*>& list, const mjKeyM
// find object in global lists given string type and name
mjCBase* mjCModel::FindObject(mjtObj type, string name) {
switch (type) {
case mjOBJ_BODY:
case mjOBJ_XBODY:
return findobject(name, bodies, ids["body"]);
case mjOBJ_JOINT:
return findobject(name, joints, ids["joint"]);
case mjOBJ_GEOM:
return findobject(name, geoms, ids["geom"]);
case mjOBJ_SITE:
return findobject(name, sites, ids["site"]);
case mjOBJ_CAMERA:
return findobject(name, cameras, ids["camera"]);
case mjOBJ_LIGHT:
return findobject(name, lights, ids["light"]);
case mjOBJ_FLEX:
return findobject(name, flexes, ids["flex"]);
case mjOBJ_MESH:
return findobject(name, meshes, ids["mesh"]);
case mjOBJ_SKIN:
return findobject(name, skins, ids["skin"]);
case mjOBJ_HFIELD:
return findobject(name, hfields, ids["hfield"]);
case mjOBJ_TEXTURE:
return findobject(name, textures, ids["texture"]);
case mjOBJ_MATERIAL:
return findobject(name, materials, ids["material"]);
case mjOBJ_PAIR:
return findobject(name, pairs, ids["pair"]);
case mjOBJ_EXCLUDE:
return findobject(name, excludes, ids["exclude"]);
case mjOBJ_EQUALITY:
return findobject(name, equalities, ids["equality"]);
case mjOBJ_TENDON:
return findobject(name, tendons, ids["tendon"]);
case mjOBJ_ACTUATOR:
return findobject(name, actuators, ids["actuator"]);
case mjOBJ_SENSOR:
return findobject(name, sensors, ids["sensor"]);
case mjOBJ_NUMERIC:
return findobject(name, numerics, ids["numeric"]);
case mjOBJ_TEXT:
return findobject(name, texts, ids["text"]);
case mjOBJ_TUPLE:
return findobject(name, tuples, ids["tuple"]);
case mjOBJ_PLUGIN:
return findobject(name, plugins, ids["plugin"]);
default:
return 0;
if (!object_lists[type]) {
return nullptr;
}
return findobject(name, *object_lists[type], ids[type]);
}
@@ -2649,37 +2533,37 @@ static void reassignid(vector<T*>& list) {
// set ids, check for repeated names
template <class T>
static void processlist(mjListKeyMap& ids, vector<T*>& list,
string defname, bool checkrepeat = true) {
mjtObj type, bool checkrepeat = true) {
// loop over list elements
for (int i=0; i<(int)list.size(); i++) {
for (size_t i=0; i < list.size(); i++) {
// check for incompatible id setting; SHOULD NOT OCCUR
if (list[i]->id!=-1 && list[i]->id!=i) {
throw mjCError(list[i], "incompatible id in %s array, position %d", defname.c_str(), i);
throw mjCError(list[i], "incompatible id in %s array, position %d", mju_type2Str(type), i);
}
// id equals position in array
list[i]->id = i;
// add to ids map
ids[defname][list[i]->name] = i;
ids[type][list[i]->name] = i;
}
// check for repeated names
if (checkrepeat) {
// created vectors with all names
vector<string> allnames;
for (int i=0; i<(int)list.size(); i++) {
for (size_t i=0; i < list.size(); i++) {
if (!list[i]->name.empty()) {
allnames.push_back(list[i]->name);
}
}
// sort and check for duplicates
if (allnames.size()>1) {
if (allnames.size() > 1) {
std::sort(allnames.begin(), allnames.end());
auto adjacent = std::adjacent_find(allnames.begin(), allnames.end());
if (adjacent!=allnames.end()) {
string msg = "repeated name '" + *adjacent + "' in " + defname;
if (adjacent != allnames.end()) {
string msg = "repeated name '" + *adjacent + "' in " + mju_type2Str(type);
throw mjCError(NULL, msg.c_str());
}
}
@@ -2813,29 +2697,11 @@ void mjCModel::TryCompile(mjModel*& m, mjData*& d, const mjVFS* vfs) {
CheckEmptyNames();
// set object ids, check for repeated names
processlist(ids, bodies, "body");
processlist(ids, joints, "joint");
processlist(ids, geoms, "geom");
processlist(ids, sites, "site");
processlist(ids, cameras, "camera");
processlist(ids, lights, "light");
processlist(ids, flexes, "flex");
processlist(ids, meshes, "mesh");
processlist(ids, skins, "skin");
processlist(ids, hfields, "hfield");
processlist(ids, textures, "texture");
processlist(ids, materials, "material");
processlist(ids, pairs, "pair");
processlist(ids, excludes, "exclude");
processlist(ids, equalities, "equality");
processlist(ids, tendons, "tendon");
processlist(ids, actuators, "actuator");
processlist(ids, sensors, "sensor");
processlist(ids, numerics, "numeric");
processlist(ids, texts, "text");
processlist(ids, tuples, "tuple");
processlist(ids, keys, "key");
processlist(ids, plugins, "plugin");
for (int i = 0; i < mjNOBJECT; i++) {
if (i != mjOBJ_XBODY && object_lists[i]) {
processlist(ids, *object_lists[i], (mjtObj) i);
}
}
// convert names into indices
IndexAssets();
+4 -2
View File
@@ -32,9 +32,8 @@ typedef enum _mjtInertiaFromGeom {
mjINERTIAFROMGEOM_AUTO // use only if inertial element is missing
} mjtInertiaFromGeom;
// TODO: convert mjListKeyMap to mjKeyMap[mjNOBJECT] by adding mjNOBJECT to mjtObj
typedef std::map<std::string, int, std::less<> > mjKeyMap;
typedef std::map<std::string, mjKeyMap, std::less<> > mjListKeyMap;
typedef std::array<mjKeyMap, mjNOBJECT> mjListKeyMap;
@@ -294,6 +293,9 @@ class mjCModel {
//------------------------ internal variables
// array of pointers to each object list (enumerated by type)
std::array<std::vector<mjCBase*>*, mjNOBJECT> object_lists;
// statistics, as computed by mj_setConst
double meaninertia_auto; // mean diagonal inertia, as computed by mj_setConst
double meanmass_auto; // mean body mass, as computed by mj_setConst
+26 -1
View File
@@ -30,16 +30,41 @@ namespace {
using ::testing::DoubleNear;
using ::testing::ElementsAre;
using ::testing::HasSubstr;
using ::testing::IsNull;
using ::testing::NotNull;
using UserDataTest = MujocoTest;
static std::vector<mjtNum> GetRow(const mjtNum* array, int ncolumn, int row) {
return std::vector<mjtNum>(array + ncolumn * row,
array + ncolumn * (row + 1));
}
// ----------------------------- test mjCModel --------------------------------
using UserCModelTest = MujocoTest;
TEST_F(UserCModelTest, RepeatedNames) {
static constexpr char xml[] = R"(
<mujoco>
<worldbody>
<body name="body1">
<joint axis="0 1 0" name="joint1"/>
<geom size="1" name="geom1"/>
<geom size="1" name="geom1"/>
</body>
</worldbody>
</mujoco>)";
std::array<char, 1024> error;
mjModel* model = LoadModelFromString(xml, error.data(), error.size());
EXPECT_THAT(model, IsNull());
EXPECT_THAT(error.data(), HasSubstr("repeated name 'geom1' in geom"));
}
// ------------- test automatic inference of nuser_xxx -------------------------
using UserDataTest = MujocoTest;
TEST_F(UserDataTest, AutoNUserBody) {
static constexpr char xml[] = R"(
<mujoco>