Implement hash map for mj_name2id using djb2 and linear probing.

PiperOrigin-RevId: 499485127
Change-Id: I5aed6549db18b92d7c8a3b7590e50a00ce13a439
This commit is contained in:
Kyle Bayes
2023-01-04 07:57:11 -08:00
committed by Copybara-Service
parent 5261cceb30
commit f1007df013
11 changed files with 404 additions and 106 deletions
+190 -81
View File
@@ -20,6 +20,7 @@
#include <mujoco/mjmodel.h>
#include "engine/engine_array_safety.h"
#include "engine/engine_core_constraint.h"
#include "engine/engine_crossplatform.h"
#include "engine/engine_io.h"
#include "engine/engine_macro.h"
#include "engine/engine_util_blas.h"
@@ -433,123 +434,231 @@ int mj_jacDifPair(const mjModel* m, const mjData* d, int* chain,
//-------------------------- name functions --------------------------------------------------------
// get number of objects and name addresses for given object type
static int _getnumadr(const mjModel* m, mjtObj type, int** padr) {
static int _getnumadr(const mjModel* m, mjtObj type, int** padr, int* mapadr) {
int num = -1;
// map address starts at the end, subtract with explicit switch fallthrough below
*mapadr = m->nnames_map;
// get address list and size for object type
switch (type) {
case mjOBJ_BODY:
case mjOBJ_XBODY:
*padr = m->name_bodyadr;
return m->nbody;
case mjOBJ_BODY:
case mjOBJ_XBODY:
*mapadr -= m->nbody;
*padr = m->name_bodyadr;
num = m->nbody;
mjFALLTHROUGH;
case mjOBJ_JOINT:
*padr = m->name_jntadr;
return m->njnt;
case mjOBJ_JOINT:
*mapadr -= m->njnt;
if (num < 0) {
*padr = m->name_jntadr;
num = m->njnt;
}
mjFALLTHROUGH;
case mjOBJ_GEOM:
*padr = m->name_geomadr;
return m->ngeom;
case mjOBJ_GEOM:
*mapadr -= m->ngeom;
if (num < 0) {
*padr = m->name_geomadr;
num = m->ngeom;
}
mjFALLTHROUGH;
case mjOBJ_SITE:
*padr = m->name_siteadr;
return m->nsite;
case mjOBJ_SITE:
*mapadr -= m->nsite;
if (num < 0) {
*padr = m->name_siteadr;
num = m->nsite;
}
mjFALLTHROUGH;
case mjOBJ_CAMERA:
*padr = m->name_camadr;
return m->ncam;
case mjOBJ_CAMERA:
*mapadr -= m->ncam;
if (num < 0) {
*padr = m->name_camadr;
num = m->ncam;
}
mjFALLTHROUGH;
case mjOBJ_LIGHT:
*padr = m->name_lightadr;
return m->nlight;
case mjOBJ_LIGHT:
*mapadr -= m->nlight;
if (num < 0) {
*padr = m->name_lightadr;
num = m->nlight;
}
mjFALLTHROUGH;
case mjOBJ_MESH:
*padr = m->name_meshadr;
return m->nmesh;
case mjOBJ_MESH:
*mapadr -= m->nmesh;
if (num < 0) {
*padr = m->name_meshadr;
num = m->nmesh;
}
mjFALLTHROUGH;
case mjOBJ_SKIN:
*padr = m->name_skinadr;
return m->nskin;
case mjOBJ_SKIN:
*mapadr -= m->nskin;
if (num < 0) {
*padr = m->name_skinadr;
num = m->nskin;
}
mjFALLTHROUGH;
case mjOBJ_HFIELD:
*padr = m->name_hfieldadr;
return m->nhfield;
case mjOBJ_HFIELD:
*mapadr -= m->nhfield;
if (num < 0) {
*padr = m->name_hfieldadr;
num = m->nhfield;
}
mjFALLTHROUGH;
case mjOBJ_TEXTURE:
*padr = m->name_texadr;
return m->ntex;
case mjOBJ_TEXTURE:
*mapadr -= m->ntex;
if (num < 0) {
*padr = m->name_texadr;
num = m->ntex;
}
mjFALLTHROUGH;
case mjOBJ_MATERIAL:
*padr = m->name_matadr;
return m->nmat;
case mjOBJ_MATERIAL:
*mapadr -= m->nmat;
if (num < 0) {
*padr = m->name_matadr;
num = m->nmat;
}
mjFALLTHROUGH;
case mjOBJ_PAIR:
*padr = m->name_pairadr;
return m->npair;
case mjOBJ_PAIR:
*mapadr -= m->npair;
if (num < 0) {
*padr = m->name_pairadr;
num = m->npair;
}
mjFALLTHROUGH;
case mjOBJ_EXCLUDE:
*padr = m->name_excludeadr;
return m->nexclude;
case mjOBJ_EXCLUDE:
*mapadr -= m->nexclude;
if (num < 0) {
*padr = m->name_excludeadr;
num = m->nexclude;
}
mjFALLTHROUGH;
case mjOBJ_EQUALITY:
*padr = m->name_eqadr;
return m->neq;
case mjOBJ_EQUALITY:
*mapadr -= m->neq;
if (num < 0) {
*padr = m->name_eqadr;
num = m->neq;
}
mjFALLTHROUGH;
case mjOBJ_TENDON:
*padr = m->name_tendonadr;
return m->ntendon;
case mjOBJ_TENDON:
*mapadr -= m->ntendon;
if (num < 0) {
*padr = m->name_tendonadr;
num = m->ntendon;
}
mjFALLTHROUGH;
case mjOBJ_ACTUATOR:
*padr = m->name_actuatoradr;
return m->nu;
case mjOBJ_ACTUATOR:
*mapadr -= m->nu;
if (num < 0) {
*padr = m->name_actuatoradr;
num = m->nu;
}
mjFALLTHROUGH;
case mjOBJ_SENSOR:
*padr = m->name_sensoradr;
return m->nsensor;
case mjOBJ_SENSOR:
*mapadr -= m->nsensor;
if (num < 0) {
*padr = m->name_sensoradr;
num = m->nsensor;
}
mjFALLTHROUGH;
case mjOBJ_NUMERIC:
*padr = m->name_numericadr;
return m->nnumeric;
case mjOBJ_NUMERIC:
*mapadr -= m->nnumeric;
if (num < 0) {
*padr = m->name_numericadr;
num = m->nnumeric;
}
mjFALLTHROUGH;
case mjOBJ_TEXT:
*padr = m->name_textadr;
return m->ntext;
case mjOBJ_TEXT:
*mapadr -= m->ntext;
if (num < 0) {
*padr = m->name_textadr;
num = m->ntext;
}
mjFALLTHROUGH;
case mjOBJ_TUPLE:
*padr = m->name_tupleadr;
return m->ntuple;
case mjOBJ_TUPLE:
*mapadr -= m->ntuple;
if (num < 0) {
*padr = m->name_tupleadr;
num = m->ntuple;
}
mjFALLTHROUGH;
case mjOBJ_KEY:
*padr = m->name_keyadr;
return m->nkey;
case mjOBJ_KEY:
*mapadr -= m->nkey;
if (num < 0) {
*padr = m->name_keyadr;
num = m->nkey;
}
mjFALLTHROUGH;
case mjOBJ_PLUGIN:
*padr = m->name_pluginadr;
return m->nplugin;
case mjOBJ_PLUGIN:
*mapadr -= m->nplugin;
if (num < 0) {
*padr = m->name_pluginadr;
num = m->nplugin;
}
mjFALLTHROUGH;
default:
*padr = 0;
return 0;
default:
if (num < 0) {
*padr = 0;
num = 0;
}
}
return num;
}
// get string hash, see http://www.cse.yorku.ca/~oz/hash.html
uint64_t mj_hashdjb2(const char* s, uint64_t n) {
uint64_t h = 5381;
int c;
while ((c = *s++)) {
h = ((h << 5) + h) + c;
}
return h % n;
}
// get id of object with specified name; -1: not found
int mj_name2id(const mjModel* m, int type, const char* name) {
int num = 0;
int mapadr;
int* adr = 0;
// get number of objects and name addresses
num = _getnumadr(m, type, &adr);
int num = _getnumadr(m, type, &adr, &mapadr);
// search
if (num) {
for (int i=0; i<num; i++) {
if (!strncmp(name, m->names+adr[i], m->nnames-adr[i])) {
return i;
}
}
}
// look up at hash address
uint64_t hash = mj_hashdjb2(name, num);
uint64_t i = hash;
do {
int j = m->names_map[mapadr + i];
if (j < 0) return -1;
if (!strncmp(name, m->names+adr[j], m->nnames-adr[j])) {
return j;
}
if (++i == num) i = 0;
} while(i != hash);
}
return -1;
}
@@ -557,11 +666,11 @@ int mj_name2id(const mjModel* m, int type, const char* name) {
// get name of object with specified id; 0: invalid type or id, or null name
const char* mj_id2name(const mjModel* m, int type, int id) {
int num = 0;
int mapadr;
int* adr = 0;
// get number of objects and name addresses
num = _getnumadr(m, type, &adr);
int num = _getnumadr(m, type, &adr, &mapadr);
if (id>=0 && id<num) {
if (m->names[adr[id]]) {