Add contact sensor.

PiperOrigin-RevId: 783011982
Change-Id: Ica56fe9d520fa1d1ee7338e09148b1a55a049912
This commit is contained in:
Yuval Tassa
2025-07-14 13:01:55 -07:00
committed by Copybara-Service
parent e441868dad
commit d0e4771c8c
25 changed files with 1343 additions and 47 deletions
+1
View File
@@ -2159,6 +2159,7 @@ static int sensorSize(mjtSensor sensor_type, int sensor_dim) {
case mjSENS_FRAMEQUAT:
return 4;
case mjSENS_CONTACT:
case mjSENS_USER:
return sensor_dim;
+211 -1
View File
@@ -26,6 +26,7 @@
#include "engine/engine_io.h"
#include "engine/engine_plugin.h"
#include "engine/engine_ray.h"
#include "engine/engine_sort.h"
#include "engine/engine_support.h"
#include "engine/engine_util_blas.h"
#include "engine/engine_util_errmem.h"
@@ -36,6 +37,26 @@
//-------------------------------- utility ---------------------------------------------------------
typedef struct {
mjtNum criterion; // criterion for partial sort
int id; // index in d->contact
int flip; // 0: don't flip the normal, 1: flip the normal
} ContactInfo;
// define ContactSelect: find the k smallest elements of a ContactInfo array
static int ContactInfoCompare(const ContactInfo* a, const ContactInfo* b, void* context) {
if (a->criterion < b->criterion) return -1;
if (a->criterion > b->criterion) return 1;
if (a->id < b->id) return -1;
if (a->id > b->id) return 1;
return 0;
}
mjPARTIAL_SORT(ContactSelect, ContactInfo, ContactInfoCompare)
// apply cutoff after each stage
static void apply_cutoff(const mjModel* m, mjData* d, mjtStage stage) {
// process sensors matching stage and having positive cutoff
@@ -98,6 +119,8 @@ static void get_xpos_xmat(const mjData* d, mjtObj type, int id, int sensor_id,
}
}
// get global quaternion of an object in mjData
static void get_xquat(const mjModel* m, const mjData* d, mjtObj type, int id, int sensor_id,
mjtNum *quat) {
@@ -123,6 +146,7 @@ static void get_xquat(const mjModel* m, const mjData* d, mjtObj type, int id, in
}
static void cam_project(mjtNum sensordata[2], const mjtNum target_xpos[3],
const mjtNum cam_xpos[3], const mjtNum cam_xmat[9],
const int cam_res[2], mjtNum cam_fovy,
@@ -216,6 +240,117 @@ static void cam_project(mjtNum sensordata[2], const mjtNum target_xpos[3],
// check if a contact body/geom matches a sensor spec (type, id)
static int checkMatch(const mjModel* m, int body, int geom, mjtObj type, int id) {
if (type == mjOBJ_UNKNOWN) return 1;
if (type == mjOBJ_SITE) return 1; // already passed site filter test
if (type == mjOBJ_GEOM) return id == geom;
if (type == mjOBJ_BODY) return id == body;
if (type == mjOBJ_XBODY) return body >= 0 && m->body_rootid[id] == m->body_rootid[body];
return 0;
}
// 0: no match
// 1: match, use contact normal
// -1: match, flip contact normal
static int matchContact(const mjModel* m, const mjData* d, int conid,
mjtObj type1, int id1, mjtObj type2, int id2) {
// no criterion: quick match
if (type1 == mjOBJ_UNKNOWN && type2 == mjOBJ_UNKNOWN) {
return 1;
}
// site filter
if (type1 == mjOBJ_SITE) {
if (!mju_insideGeom(d->site_xpos + 3 * id1, d->site_xmat + 9 * id1,
m->site_size + 3 * id1, m->site_type[id1], d->contact[conid].pos)) {
return 0;
}
}
// get geom, body ids
int geom1 = d->contact[conid].geom[0];
int geom2 = d->contact[conid].geom[1];
int body1 = geom1 >= 0 ? m->geom_bodyid[geom1] : -1;
int body2 = geom2 >= 0 ? m->geom_bodyid[geom2] : -1;
// check match of sensor objects with contact objects
int match11 = checkMatch(m, body1, geom1, type1, id1);
int match12 = checkMatch(m, body2, geom2, type1, id1);
int match21 = checkMatch(m, body1, geom1, type2, id2);
int match22 = checkMatch(m, body2, geom2, type2, id2);
// if a sensor object is specified, it must be involved in the contact
if (!match11 && !match12) return 0;
if (!match21 && !match22) return 0;
// determine direction
if (type1 != mjOBJ_UNKNOWN && type2 != mjOBJ_UNKNOWN) {
// both obj1 and obj2 specified: direction depends on order
int order_regular = match11 && match22;
int order_reverse = match12 && match21;
if (order_regular && !order_reverse) return 1;
if (order_reverse && !order_regular) return -1;
if (order_regular && order_reverse) return 1; // ambiguous, return 1
} else if (type1 != mjOBJ_UNKNOWN) {
// only obj1 specified: normal points away from obj1
return match11 ? 1 : -1;
} else if (type2 != mjOBJ_UNKNOWN) {
// only obj2 specified: normal points towards obj2
return match22 ? 1 : -1;
}
// should not occur, all conditions are covered above
return 0;
}
// fill in output data for contact sensor for all fields
// if flg_flip > 0, normal/tangent rotate 180 about frame[2]
// force/torque flip-z s.t. force is equal-and-opposite in new contact frame
static void copySensorData(const mjModel* m, const mjData* d,
mjtNum* data[mjNCONDATA], int id, int flg_flip, int nfound) {
// found flag
if (data[mjCONDATA_FOUND]) *data[mjCONDATA_FOUND] = nfound;
// contact force and torque
if (data[mjCONDATA_FORCE] || data[mjCONDATA_TORQUE]) {
mjtNum forcetorque[6];
mj_contactForce(m, d, id, forcetorque);
if (data[mjCONDATA_FORCE]) {
mju_copy3(data[mjCONDATA_FORCE], forcetorque);
if (flg_flip) data[mjCONDATA_FORCE][2] *= -1;
}
if (data[mjCONDATA_TORQUE]) {
mju_copy3(data[mjCONDATA_TORQUE], forcetorque+3);
if (flg_flip) data[mjCONDATA_TORQUE][2] *= -1;
}
}
// contact penetration distance
if (data[mjCONDATA_DIST]) {
*data[mjCONDATA_DIST] = d->contact[id].dist;
}
// contact position
if (data[mjCONDATA_POS]) {
mju_copy3(data[mjCONDATA_POS], d->contact[id].pos);
}
// contact normal
if (data[mjCONDATA_NORMAL]) {
mju_copy3(data[mjCONDATA_NORMAL], d->contact[id].frame);
if (flg_flip) mju_scl3(data[mjCONDATA_NORMAL], data[mjCONDATA_NORMAL], -1);
}
// contact first tangent
if (data[mjCONDATA_TANGENT]) {
mju_copy3(data[mjCONDATA_TANGENT], d->contact[id].frame+3);
if (flg_flip) mju_scl3(data[mjCONDATA_TANGENT], data[mjCONDATA_TANGENT], -1);
}
}
//-------------------------------- sensor ----------------------------------------------------------
// position-dependent sensors
@@ -709,7 +844,7 @@ void mj_sensorAcc(const mjModel* m, mjData* d) {
int rootid, bodyid, objtype, objid, adr, nusersensor = 0;
int ne = d->ne, nf = d->nf, nefc = d->nefc, nu = m->nu;
mjtNum tmp[6], conforce[6], conray[3], frc;
mjContact* con;
const mjContact* con;
// disabled sensors: return
if (mjDISABLED(mjDSBL_SENSOR)) {
@@ -792,6 +927,81 @@ void mj_sensorAcc(const mjModel* m, mjData* d) {
}
break;
case mjSENS_CONTACT: // contact
{
// prepare sizes and indices, check consistency
int dataspec = m->sensor_intprm[i*mjNSENS];
int size = mju_condataSize(dataspec); // size of each slot
int dim = m->sensor_dim[i]; // total sensor array dimension
int num = dim / size; // number of slots
int reftype = m->sensor_reftype[i];
int refid = m->sensor_refid[i];
int reduce = m->sensor_intprm[i*mjNSENS+1];
// clear all outputs, prepare data pointers
mjtNum* ptr = d->sensordata + adr;
mju_zero(ptr, dim);
mjtNum* data[mjNCONDATA] = {NULL};
for (int j=0; j < mjNCONDATA; j++) {
if (dataspec & (1 << j)) {
data[j] = ptr;
ptr += mjCONDATA_SIZE[j];
}
}
// prepare for matching loop
int nmatch = 0;
mj_markStack(d);
ContactInfo *match = mjSTACKALLOC(d, d->ncon, ContactInfo);
// find matching contacts
for (int j=0; j < d->ncon; j++) {
// check match condition
int match_j = matchContact(m, d, j, objtype, objid, reftype, refid);
if (!match_j) {
continue;
}
// save id and flip flag
match[nmatch].id = j;
match[nmatch].flip = match_j < 0;
// save sorting criterion, if required
if (reduce) {
if (reduce == 1) {
match[nmatch].criterion = d->contact[j].dist;
} else {
mjtNum forcetorque[6];
mj_contactForce(m, d, j, forcetorque);
match[nmatch].criterion = -mju_dot3(forcetorque, forcetorque);
}
}
// increment number of matching contacts
nmatch++;
}
// number of slots to be filled
int nslot = mjMIN(num, nmatch);
// partial sort to get bottom nslot contacts given reduction criterion
if (reduce) {
ContactInfo *heap = mjSTACKALLOC(d, nslot, ContactInfo);
ContactSelect(match, heap, nmatch, nslot, NULL);
}
// copy data into slots, increment pointers
for (int j=0; j < nslot; j++) {
copySensorData(m, d, data, match[j].id, match[j].flip, nmatch);
for (int k=0; k < mjNCONDATA; k++) {
if (data[k]) data[k] += size;
}
}
mj_freeStack(d);
}
break;
case mjSENS_ACCELEROMETER: // accelerometer
// tmp = site acceleration, in site frame
mj_objectAcceleration(m, d, mjOBJ_SITE, objid, tmp, 1);
+24
View File
@@ -97,6 +97,17 @@ const char* mjTIMERSTRING[mjNTIMER]= {
};
// size of contact data fields
const int mjCONDATA_SIZE[mjNCONDATA] = {
1, // mjCONDATA_FOUND
3, // mjCONDATA_FORCE
3, // mjCONDATA_TORQUE
1, // mjCONDATA_DIST
3, // mjCONDATA_POS
3, // mjCONDATA_NORMAL
3 // mjCONDATA_TANGENT
};
//-------------------------- get/set state ---------------------------------------------------------
@@ -1576,3 +1587,16 @@ const char* mj_versionString(void) {
static const char versionstring[] = mjVERSIONSTRING;
return versionstring;
}
// return total size of data in a contact sensor bitfield specification
int mju_condataSize(int dataspec) {
int size = 0;
for (int i=0; i < mjNCONDATA; i++) {
if (dataspec & (1 << i)) {
size += mjCONDATA_SIZE[i];
}
}
return size;
}
+7
View File
@@ -29,6 +29,9 @@ MJAPI extern const char* mjDISABLESTRING[mjNDISABLE];
MJAPI extern const char* mjENABLESTRING[mjNENABLE];
MJAPI extern const char* mjTIMERSTRING[mjNTIMER];
// arrays
MJAPI extern const int mjCONDATA_SIZE[mjNCONDATA]; // TODO(tassa): expose in public header?
//-------------------------- get/set state ---------------------------------------------------------
@@ -199,6 +202,10 @@ MJAPI int mj_version(void);
// current version of MuJoCo as a null-terminated string
MJAPI const char* mj_versionString(void);
// return total size of data fields in a contact sensor bitfield specification
MJAPI int mju_condataSize(int dataSpec);
#ifdef __cplusplus
}
#endif
+58 -3
View File
@@ -37,6 +37,7 @@
#include "lodepng.h"
#include "cc/array_safety.h"
#include "engine/engine_passive.h"
#include "engine/engine_support.h"
#include <mujoco/mjspec.h>
#include <mujoco/mujoco.h>
#include "user/user_api.h"
@@ -6686,6 +6687,7 @@ void mjCSensor::ResolveReferences(const mjCModel* m) {
type != mjSENS_E_KINETIC &&
type != mjSENS_CLOCK &&
type != mjSENS_PLUGIN &&
type != mjSENS_CONTACT &&
type != mjSENS_USER) {
throw mjCError(this, "invalid type in sensor");
}
@@ -7043,6 +7045,60 @@ void mjCSensor::Compile(void) {
}
break;
case mjSENS_CONTACT:
// check first matching criterion
if (objtype != mjOBJ_SITE &&
objtype != mjOBJ_BODY &&
objtype != mjOBJ_XBODY &&
objtype != mjOBJ_GEOM &&
objtype != mjOBJ_UNKNOWN) {
throw mjCError(this, "first matching criterion: if set, must be (x)body, geom or site");
}
// check that subtree1 is a full tree
if (objtype == mjOBJ_XBODY && static_cast<mjCBody*>(obj)->GetParent()->id != 0) {
throw mjCError(this, "subtree1 must be a child of the world");
}
// check second matching criterion
if (reftype != mjOBJ_BODY &&
reftype != mjOBJ_XBODY &&
reftype != mjOBJ_GEOM &&
reftype != mjOBJ_UNKNOWN) {
throw mjCError(this, "second matching criterion: if set, must be (x)body or geom");
}
// check that subtree2 is a full tree
if (reftype == mjOBJ_XBODY && static_cast<mjCBody*>(ref)->GetParent()->id != 0) {
throw mjCError(this, "subtree2 must be a child of the world");
}
// check for non-positive dim
if (dim <= 0) {
throw mjCError(this, "dim must be positive in sensor (got %d)", "", dim);
}
// check for dim correctness
if (dim % mju_condataSize(intprm[0]) != 0) {
throw mjCError(this, "dim %d does not match data spec", "", dim);
}
// check for reduce correctness
if (intprm[1] < 0 || intprm[1] > 3) {
throw mjCError(this, "unknown reduction criterion. got %d, "
"expected one of {0, 1, 2, 3}", "", intprm[1]);
}
// netforce not yet implemented
if (intprm[1] == 3) {
throw mjCError(this, "netforce reduction is not yet implemented\n"
"please contact the developers if you need this feature");
}
needstage = mjSTAGE_ACC;
datatype = mjDATATYPE_REAL;
break;
case mjSENS_E_POTENTIAL:
case mjSENS_E_KINETIC:
case mjSENS_CLOCK:
@@ -7054,13 +7110,12 @@ void mjCSensor::Compile(void) {
case mjSENS_USER:
// check for negative dim
if (dim < 0) {
throw mjCError(this, "sensor dim must be positive in sensor");
throw mjCError(this, "sensor dim must be non-negative in sensor");
}
// make sure dim is consistent with datatype
if (datatype == mjDATATYPE_AXIS && dim != 3) {
throw mjCError(this,
"datatype AXIS requires dim=3 in sensor");
throw mjCError(this, "datatype AXIS requires dim=3 in sensor");
}
if (datatype == mjDATATYPE_QUATERNION && dim != 4) {
throw mjCError(this, "datatype QUATERNION requires dim=4 in sensor");
+3
View File
@@ -43,6 +43,7 @@ extern const int gain_sz;
extern const int bias_sz;
extern const int stage_sz;
extern const int datatype_sz;
extern const int reduce_sz;
extern const mjMap angle_map[];
extern const mjMap enable_map[];
extern const mjMap bool_map[];
@@ -70,6 +71,8 @@ extern const mjMap gain_map[];
extern const mjMap bias_map[];
extern const mjMap stage_map[];
extern const mjMap datatype_map[];
extern const mjMap condata_map[];
extern const mjMap reduce_map[];
extern const mjMap meshtype_map[];
extern const mjMap meshinertia_map[];
extern const mjMap flexself_map[];
+95
View File
@@ -34,6 +34,7 @@
#include <mujoco/mjtnum.h>
#include <mujoco/mjvisualize.h>
#include "engine/engine_plugin.h"
#include "engine/engine_support.h"
#include "engine/engine_util_errmem.h"
#include "engine/engine_util_misc.h"
#include <mujoco/mjspec.h>
@@ -481,6 +482,8 @@ const char* MJCF[nMJCF][mjXATTRNUM] = {
{"distance", "*", "8", "name", "geom1", "geom2", "body1", "body2", "cutoff", "noise", "user"},
{"normal", "*", "8", "name", "geom1", "geom2", "body1", "body2", "cutoff", "noise", "user"},
{"fromto", "*", "8", "name", "geom1", "geom2", "body1", "body2", "cutoff", "noise", "user"},
{"contact", "*", "12", "name", "geom1", "geom2", "body1", "body2", "subtree1", "subtree2", "site",
"num", "data", "reduce", "cutoff", "noise", "user"},
{"e_potential", "*", "4", "name", "cutoff", "noise", "user"},
{"e_kinetic", "*", "4", "name", "cutoff", "noise", "user"},
{"clock", "*", "4", "name", "cutoff", "noise", "user"},
@@ -606,6 +609,7 @@ const mjMap texrole_map[texrole_sz] = {
{"orm", mjTEXROLE_ORM},
};
// integrator type
const int integrator_sz = 4;
const mjMap integrator_map[integrator_sz] = {
@@ -615,6 +619,7 @@ const mjMap integrator_map[integrator_sz] = {
{"implicitfast", mjINT_IMPLICITFAST}
};
// cone type
const int cone_sz = 2;
const mjMap cone_map[cone_sz] = {
@@ -743,6 +748,28 @@ const mjMap datatype_map[datatype_sz] = {
};
// contact data type
const mjMap condata_map[mjNCONDATA] = {
{"found", mjCONDATA_FOUND},
{"force", mjCONDATA_FORCE},
{"torque", mjCONDATA_TORQUE},
{"dist", mjCONDATA_DIST},
{"pos", mjCONDATA_POS},
{"normal", mjCONDATA_NORMAL},
{"tangent", mjCONDATA_TANGENT}
};
// contact reduction type
const int reduce_sz = 4;
const mjMap reduce_map[reduce_sz] = {
{"none", 0},
{"mindist", 1},
{"maxforce", 2},
{"netforce", 3}
};
// LR mode
const int lrmode_sz = 4;
const mjMap lrmode_map[lrmode_sz] = {
@@ -4132,6 +4159,74 @@ void mjXReader::Sensor(XMLElement* section) {
}
}
// sensor for contacts; attached to geoms or bodies or a site
else if (type == "contact") {
// first matching criterion
bool has_site = ReadAttrTxt(elem, "site", objname);
bool has_body1 = ReadAttrTxt(elem, "body1", objname);
bool has_subtree1 = ReadAttrTxt(elem, "subtree1", objname);
bool has_geom1 = ReadAttrTxt(elem, "geom1", objname);
if (has_site + has_body1 + has_subtree1 + has_geom1 > 1) {
throw mjXError(elem, "at most one of (geom1, body1, subtree1, site) can be specified");
}
if (has_site) { sensor->objtype = mjOBJ_SITE; }
else if (has_body1) { sensor->objtype = mjOBJ_BODY; }
else if (has_subtree1) { sensor->objtype = mjOBJ_XBODY; }
else if (has_geom1) { sensor->objtype = mjOBJ_GEOM; }
else { sensor->objtype = mjOBJ_UNKNOWN; }
// second matching criterion
bool has_body2 = ReadAttrTxt(elem, "body2", refname);
bool has_subtree2 = ReadAttrTxt(elem, "subtree2", refname);
bool has_geom2 = ReadAttrTxt(elem, "geom2", refname);
if (has_body2 + has_subtree2 + has_geom2 > 1) {
throw mjXError(elem, "at most one of (geom2, body2, subtree2) can be specified");
}
if (has_body2) { sensor->reftype = mjOBJ_BODY; }
else if (has_subtree2) { sensor->reftype = mjOBJ_XBODY; }
else if (has_geom2) { sensor->reftype = mjOBJ_GEOM; }
else { sensor->reftype = mjOBJ_UNKNOWN; }
// process data specification (intprm[0])
int dataspec = 1 << mjCONDATA_FOUND;
std::vector<int> condata(mjNCONDATA);
int nkeys = MapValues(elem, "data", condata.data(), condata_map, mjNCONDATA);
if (nkeys) {
dataspec = 1 << condata[0];
// check ordering while adding bits to dataspec
for (int i = 1; i < nkeys; ++i) {
if (condata[i] <= condata[i-1]) {
std::string correct_order;
for (int j = 0; j < mjNCONDATA; ++j) {
correct_order += condata_map[j].key;
if (j < mjNCONDATA - 1) correct_order += ", ";
}
throw mjXError(elem, "data attributes must be in order: %s", correct_order.c_str());
}
dataspec |= 1 << condata[i];
}
}
sensor->intprm[0] = dataspec;
// number of contacts, sensor dim
sensor->dim = 1;
ReadAttrInt(elem, "num", &sensor->dim);
if (sensor->dim <= 0) {
throw mjXError(elem, "'num' must be positive in sensor");
}
sensor->dim *= mju_condataSize(dataspec);
// reduction type (intprm[1])
sensor->intprm[1] = 0;
if (MapValue(elem, "reduce", &n, reduce_map, reduce_sz)) {
sensor->intprm[1] = n;
}
// sensor type
sensor->type = mjSENS_CONTACT;
}
// global sensors
else if (type == "e_potential") {
sensor->type = mjSENS_E_POTENTIAL;
+1 -1
View File
@@ -101,7 +101,7 @@ class mjXReader : public mjXBase {
};
// MJCF schema
#define nMJCF 239
#define nMJCF 240
extern const char* MJCF[nMJCF][mjXATTRNUM];
#endif // MUJOCO_SRC_XML_XML_NATIVE_READER_H_
+34 -2
View File
@@ -28,6 +28,7 @@
#include <mujoco/mujoco.h>
#include "engine/engine_io.h"
#include "engine/engine_plugin.h"
#include "engine/engine_support.h"
#include "engine/engine_util_errmem.h"
#include "engine/engine_util_misc.h"
#include "user/user_model.h"
@@ -76,7 +77,7 @@ static string WriteDoc(XMLDocument& doc, char *error, size_t error_sz) {
// top level sections
std::array<string, 17> sections = {
"<actuator", "<asset", "<compiler", "<contact", "<custom",
"<actuator", "<asset", "<compiler", "<contact>", "<custom",
"<default>", "<deformable", "<equality", "<extension", "<keyframe",
"<option", "<sensor", "<size", "<statistic", "<tendon",
"<visual", "<worldbody"};
@@ -2205,7 +2206,38 @@ void mjXWriter::Sensor(XMLElement* root) {
WriteAttrTxt(elem, sensor->objtype == mjOBJ_BODY ? "body1" : "geom1", sensor->get_objname());
WriteAttrTxt(elem, sensor->reftype == mjOBJ_BODY ? "body2" : "geom2", sensor->get_refname());
break;
case mjSENS_CONTACT:
{
elem = InsertEnd(section, "contact");
if (sensor->objtype == mjOBJ_BODY) {
WriteAttrTxt(elem, "body1", sensor->get_objname());
} else if (sensor->objtype == mjOBJ_XBODY) {
WriteAttrTxt(elem, "subtree1", sensor->get_objname());
} else if (sensor->objtype == mjOBJ_GEOM) {
WriteAttrTxt(elem, "geom1", sensor->get_objname());
} else if (sensor->objtype == mjOBJ_SITE) {
WriteAttrTxt(elem, "site", sensor->get_objname());
}
if (sensor->reftype == mjOBJ_BODY) {
WriteAttrTxt(elem, "body2", sensor->get_refname());
} else if (sensor->reftype == mjOBJ_XBODY) {
WriteAttrTxt(elem, "subtree2", sensor->get_refname());
} else if (sensor->reftype == mjOBJ_GEOM) {
WriteAttrTxt(elem, "geom2", sensor->get_refname());
}
int dataspec = sensor->intprm[0];
WriteAttrInt(elem, "num", sensor->dim / mju_condataSize(dataspec), 1);
int data[mjNCONDATA];
int ndata = 0;
for (int i=0; i < mjNCONDATA; i++) {
if (dataspec & (1 << i)) {
data[ndata++] = i;
}
}
WriteAttrKeys(elem, "data", condata_map, mjNCONDATA, data, ndata, 0);
WriteAttrKey(elem, "reduce", reduce_map, reduce_sz, sensor->intprm[1], 0);
}
break;
// global sensors
case mjSENS_E_POTENTIAL:
elem = InsertEnd(section, "potential");
+54 -1
View File
@@ -792,7 +792,7 @@ XMLElement* mjXUtil::FindSubElem(XMLElement* elem, std::string name, bool requir
// find attribute, translate key, return int value
// find attribute, translate key into data, return true if found
bool mjXUtil::MapValue(XMLElement* elem, const char* attr, int* data,
const mjMap* map, int mapSz, bool required) {
// get attribute text
@@ -814,6 +814,42 @@ bool mjXUtil::MapValue(XMLElement* elem, const char* attr, int* data,
// find attribute, translate unique space-separated keys to data, return number of keys found
int mjXUtil::MapValues(XMLElement* elem, const char* attr, int* data,
const mjMap* map, int mapSz, bool required) {
// get attribute text
auto maybe_text = ReadAttrStr(elem, attr, required);
if (!maybe_text.has_value()) {
return 0;
}
std::string text = maybe_text.value();
std::istringstream strm(text);
std::string key;
std::set<std::string> found_keys;
int count = 0;
while (strm >> key) {
if (found_keys.count(key)) {
throw mjXError(elem, "duplicate keyword: '%s'");
return 0;
}
int value = FindKey(map, mapSz, key);
if (value == -1) {
throw mjXError(elem, "invalid keyword: '%s'");
return 0;
}
found_keys.insert(key);
data[count++] = value;
}
return count;
}
//---------------------------------- write functions -----------------------------------------------
// check if double is int
@@ -970,3 +1006,20 @@ void mjXUtil::WriteAttrKey(XMLElement* elem, std::string name,
WriteAttrTxt(elem, name, FindValue(map, mapsz, data));
}
// write attribute- space-separated keywords
void mjXUtil::WriteAttrKeys(XMLElement* elem, std::string name, const mjMap* map,
int mapsz, int* data, int ndata, int def) {
// skip default
if (ndata == 1 && data[0] == def) {
return;
}
std::string text = FindValue(map, mapsz, data[0]);
for (int i = 1; i < ndata; ++i) {
text += " " + FindValue(map, mapsz, data[i]);
}
WriteAttrTxt(elem, name, text);
}
+8
View File
@@ -183,6 +183,10 @@ class mjXUtil {
static bool MapValue(tinyxml2::XMLElement* elem, const char* attr, int* data,
const mjMap* map, int mapSz, bool required = false);
// find attribute, translate unique space-separated keys to data, return number of keys found
static int MapValues(tinyxml2::XMLElement* elem, const char* attr, int* data,
const mjMap* map, int mapSz, bool required = false);
// write attribute- any type
template<typename T>
static void WriteAttr(tinyxml2::XMLElement* elem, std::string name, int n, const T* data,
@@ -204,6 +208,10 @@ class mjXUtil {
static void WriteAttrKey(tinyxml2::XMLElement* elem, std::string name,
const mjMap* map, int mapsz, int data, int def = -12345);
// write attribute- space-separated keywords
static void WriteAttrKeys(XMLElement* elem, std::string name, const mjMap* map,
int mapsz, int* data, int ndata, int def = -12345);
private:
template<typename T>
static bool ReadAttrValues(tinyxml2::XMLElement* elem, const char* attr,