Added mjData.eq_active user input variable, for enabling/disabling the state of equality constraints.

Renamed `mjModel.eq_active` to `mjModel.eq_active0`, which now has the semantic of "initial value of `mjData.eq_active`".

Fixes #876.

PiperOrigin-RevId: 570410643
Change-Id: Id03171e751377c7cc453f143abee64239ee2e2ed
This commit is contained in:
Yuval Tassa
2023-10-03 09:29:45 -07:00
committed by Copybara-Service
parent 177aec68dc
commit ee78b8f76b
19 changed files with 180 additions and 72 deletions
+2 -2
View File
@@ -521,7 +521,7 @@ void mj_instantiateEquality(const mjModel* m, mjData* d) {
// find active equality constraints
for (int i=0; i < m->neq; i++) {
if (m->eq_active[i]) {
if (d->eq_active[i]) {
// get constraint data
data = m->eq_data + mjNEQDATA*i;
id[0] = m->eq_obj1id[i];
@@ -1481,7 +1481,7 @@ static inline int mj_ne(const mjModel* m, mjData* d, int* nnz) {
// find active equality constraints
for (int i=0; i < neq; i++) {
if (m->eq_active[i]) {
if (d->eq_active[i]) {
id[0] = m->eq_obj1id[i];
id[1] = m->eq_obj2id[i];
size = 0;
+1
View File
@@ -1549,6 +1549,7 @@ static void _resetData(const mjModel* m, mjData* d, unsigned char debug_value) {
mju_zero(d->qvel, m->nv);
mju_zero(d->act, m->na);
mju_zero(d->ctrl, m->nu);
for (int i=0; i<m->neq; i++) d->eq_active[i] = m->eq_active0[i];
mju_zero(d->qfrc_applied, m->nv);
mju_zero(d->xfrc_applied, 6*m->nbody);
mju_zero(d->qacc, m->nv);
+7
View File
@@ -880,6 +880,13 @@ void mj_printFormattedData(const mjModel* m, mjData* d, const char* filename,
printArray("CTRL", m->nu, 1, d->ctrl, fp, float_format);
printArray("QFRC_APPLIED", m->nv, 1, d->qfrc_applied, fp, float_format);
printArray("XFRC_APPLIED", m->nbody, 6, d->xfrc_applied, fp, float_format);
if (m->neq) {
fprintf(fp, NAME_FORMAT, "EQ_ACTIVE");
for (int c=0; c < m->neq; c++) {
fprintf(fp, " %d", d->eq_active[c]);
}
fprintf(fp, "\n\n");
}
printArray("MOCAP_POS", m->nmocap, 3, d->mocap_pos, fp, float_format);
printArray("MOCAP_QUAT", m->nmocap, 4, d->mocap_quat, fp, float_format);
printArray("QACC", m->nv, 1, d->qacc, fp, float_format);
+31 -6
View File
@@ -105,6 +105,7 @@ static inline int mj_stateElemSize(const mjModel* m, mjtState spec) {
case mjSTATE_CTRL: return m->nu;
case mjSTATE_QFRC_APPLIED: return m->nv;
case mjSTATE_XFRC_APPLIED: return 6*m->nbody;
case mjSTATE_EQ_ACTIVE: return m->neq; // mjtByte, stored as mjtNum in state vector
case mjSTATE_MOCAP_POS: return 3*m->nmocap;
case mjSTATE_MOCAP_QUAT: return 4*m->nmocap;
case mjSTATE_USERDATA: return m->nuserdata;
@@ -176,9 +177,21 @@ void mj_getState(const mjModel* m, const mjData* d, mjtNum* state, unsigned int
mjtState element = 1<<i;
if (element & spec) {
int size = mj_stateElemSize(m, element);
const mjtNum* ptr = mj_stateElemConstPtr(m, d, element);
mju_copy(state + adr, ptr, size);
adr += size;
// special handling of eq_active (mjtByte)
if (element == mjSTATE_EQ_ACTIVE) {
int neq = m->neq;
for (int j=0; j < neq; j++) {
state[adr++] = d->eq_active[j];
}
}
// regular state components (mjtNum)
else {
const mjtNum* ptr = mj_stateElemConstPtr(m, d, element);
mju_copy(state + adr, ptr, size);
adr += size;
}
}
}
}
@@ -196,9 +209,21 @@ void mj_setState(const mjModel* m, mjData* d, const mjtNum* state, unsigned int
mjtState element = 1<<i;
if (element & spec) {
int size = mj_stateElemSize(m, element);
mjtNum* ptr = mj_stateElemPtr(m, d, element);
mju_copy(ptr, state + adr, size);
adr += size;
// special handling of eq_active (mjtByte)
if (element == mjSTATE_EQ_ACTIVE) {
int neq = m->neq;
for (int j=0; j < neq; j++) {
d->eq_active[j] = state[adr++];
}
}
// regular state components (mjtNum)
else {
mjtNum* ptr = mj_stateElemPtr(m, d, element);
mju_copy(ptr, state + adr, size);
adr += size;
}
}
}
}
+1 -1
View File
@@ -1908,7 +1908,7 @@ void mjv_addGeoms(const mjModel* m, mjData* d, const mjvOption* vopt,
if (vopt->flags[mjVIS_CONSTRAINT] && (category & catmask) && m->neq) {
// connect or weld
for (int i=0; i < m->neq; i++) {
if (m->eq_active[i] && (m->eq_type[i] == mjEQ_CONNECT || m->eq_type[i] == mjEQ_WELD)) {
if (d->eq_active[i] && (m->eq_type[i] == mjEQ_CONNECT || m->eq_type[i] == mjEQ_WELD)) {
// compute endpoints in global coordinates
int j = m->eq_obj1id[i], k = m->eq_obj2id[i];
mju_rotVecMat(vec, m->eq_data+mjNEQDATA*i+3*(m->eq_type[i] == mjEQ_WELD), d->xmat+9*j);