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