From 9f6be3135ecc077700ab69d60f0826bee43f24a9 Mon Sep 17 00:00:00 2001 From: Kyle Bayes Date: Thu, 14 Dec 2023 11:11:00 -0800 Subject: [PATCH] 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 --- src/user/user_model.cc | 234 ++++++++--------------------------- src/user/user_model.h | 6 +- test/user/user_model_test.cc | 27 +++- 3 files changed, 80 insertions(+), 187 deletions(-) diff --git a/src/user/user_model.cc b/src/user/user_model.cc index c1c0300e..a4f5a2dd 100644 --- a/src/user/user_model.cc +++ b/src/user/user_model.cc @@ -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*) &bodies; + object_lists[mjOBJ_XBODY] = (std::vector*) &bodies; + object_lists[mjOBJ_JOINT] = (std::vector*) &joints; + object_lists[mjOBJ_GEOM] = (std::vector*) &geoms; + object_lists[mjOBJ_SITE] = (std::vector*) &sites; + object_lists[mjOBJ_CAMERA] = (std::vector*) &cameras; + object_lists[mjOBJ_LIGHT] = (std::vector*) &lights; + object_lists[mjOBJ_FLEX] = (std::vector*) &flexes; + object_lists[mjOBJ_MESH] = (std::vector*) &meshes; + object_lists[mjOBJ_SKIN] = (std::vector*) &skins; + object_lists[mjOBJ_HFIELD] = (std::vector*) &hfields; + object_lists[mjOBJ_TEXTURE] = (std::vector*) &textures; + object_lists[mjOBJ_MATERIAL] = (std::vector*) &materials; + object_lists[mjOBJ_PAIR] = (std::vector*) &pairs; + object_lists[mjOBJ_EXCLUDE] = (std::vector*) &excludes; + object_lists[mjOBJ_EQUALITY] = (std::vector*) &equalities; + object_lists[mjOBJ_TENDON] = (std::vector*) &tendons; + object_lists[mjOBJ_ACTUATOR] = (std::vector*) &actuators; + object_lists[mjOBJ_SENSOR] = (std::vector*) &sensors; + object_lists[mjOBJ_NUMERIC] = (std::vector*) &numerics; + object_lists[mjOBJ_TEXT] = (std::vector*) &texts; + object_lists[mjOBJ_TUPLE] = (std::vector*) &tuples; + object_lists[mjOBJ_KEY] = (std::vector*) &keys; + object_lists[mjOBJ_PLUGIN] = (std::vector*) &plugins; } @@ -213,7 +241,6 @@ mjCModel::~mjCModel() { for (int i=0; isize(); } // get pointer to specified object mjCBase* mjCModel::GetObject(mjtObj type, int id) { - 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& 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& list) { // set ids, check for repeated names template static void processlist(mjListKeyMap& ids, vector& 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 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(); diff --git a/src/user/user_model.h b/src/user/user_model.h index 85ae8ec9..222cbccf 100644 --- a/src/user/user_model.h +++ b/src/user/user_model.h @@ -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 > mjKeyMap; -typedef std::map > mjListKeyMap; +typedef std::array mjListKeyMap; @@ -294,6 +293,9 @@ class mjCModel { //------------------------ internal variables + // array of pointers to each object list (enumerated by type) + std::array*, 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 diff --git a/test/user/user_model_test.cc b/test/user/user_model_test.cc index db1059ef..2079b1d1 100644 --- a/test/user/user_model_test.cc +++ b/test/user/user_model_test.cc @@ -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 GetRow(const mjtNum* array, int ncolumn, int row) { return std::vector(array + ncolumn * row, array + ncolumn * (row + 1)); } +// ----------------------------- test mjCModel -------------------------------- + +using UserCModelTest = MujocoTest; + +TEST_F(UserCModelTest, RepeatedNames) { + static constexpr char xml[] = R"( + + + + + + + + + )"; + + std::array 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"(