Improve nJmom by counting number of dofs for slidercrank and site transmissions

PiperOrigin-RevId: 875233737
Change-Id: I2cf3afad7e29e176a4f5a917dd35cfad99b2b0a3
This commit is contained in:
Taylor Howell
2026-02-25 10:40:52 -08:00
committed by Copybara-Service
parent b907ffedaa
commit 53fadd9d63
5 changed files with 51 additions and 34 deletions
+1 -1
View File
@@ -1598,7 +1598,7 @@ static int mj_jacDifPairCount(const mjModel* m, int* chain,
if (m->body_simple[b1] && m->body_simple[b2]) {
return mj_mergeChainSimple(m, chain, b1, b2);
}
return mj_mergeChain(m, chain, b1, b2);
return mj_mergeChain(m, chain, b1, b2, 0);
}
return m->nv;
+26 -24
View File
@@ -1237,7 +1237,7 @@ void mj_transmission(const mjModel* m, mjData* d) {
int *chain;
// define stack variables required for site transmission, don't allocate
mjtNum *jacref = NULL, *moment_tmp = NULL;
mjtNum *jacref = NULL, *moment_row = NULL;
int sleep_filter = mjENABLED(mjENBL_SLEEP) && d->nv_awake < nv;
@@ -1379,27 +1379,26 @@ void mj_transmission(const mjModel* m, mjData* d) {
mj_jacSite(m, d, jac, 0, id);
mju_subFrom(jac, jacS, 3*nv);
moment_row = mjSTACKALLOC(d, nv, mjtNum);
// clear moment
mju_zero(moment + adr, nv);
mju_zero(moment_row, nv);
// apply chain rule
for (int j=0; j < nv; j++) {
for (int k=0; k < 3; k++) {
moment[adr+j] += dlda[k]*jacA[k*nv+j] + dldv[k]*jac[k*nv+j];
moment_row[j] += dlda[k]*jacA[k*nv+j] + dldv[k]*jac[k*nv+j];
}
}
// scale by gear ratio
length[i] *= gear[0];
for (int j = 0; j < nv; j++) {
moment[adr+j] *= gear[0];
}
// sparsity (compress)
nnz = 0;
for (int j = 0; j < nv; j++) {
if (moment[adr+j]) {
moment[adr+nnz] = moment[adr+j];
if (moment_row[j]) {
moment[adr+nnz] = moment_row[j] * gear[0];
colind[adr+nnz] = j;
nnz++;
}
@@ -1429,6 +1428,8 @@ void mj_transmission(const mjModel* m, mjData* d) {
// clear length
length[i] = 0;
if (!moment_row) moment_row = mjSTACKALLOC(d, nv, mjtNum);
// reference site undefined
if (m->actuator_trnid[2*i+1] == -1) {
// wrench: gear expressed in global frame
@@ -1437,9 +1438,9 @@ void mj_transmission(const mjModel* m, mjData* d) {
mji_mulMatVec3(wrench+3, d->site_xmat+9*id, gear+3); // rotation
// moment: global Jacobian projected on wrench
mju_mulMatTVec(moment+adr, jac, wrench, 3, nv); // translation
mju_mulMatTVec(moment_row, jac, wrench, 3, nv); // translation
mju_mulMatTVec(jac, jacS, wrench+3, 3, nv); // rotation
mju_addTo(moment+adr, jac, nv); // add the two
mju_addTo(moment_row, jac, nv); // add the two
}
// reference site defined
@@ -1476,7 +1477,7 @@ void mj_transmission(const mjModel* m, mjData* d) {
}
// clear moment
mju_zero(moment+adr, nv);
mju_zero(moment_row, nv);
// translational transmission
if (!mju_isZero(gear, 3)) {
@@ -1508,7 +1509,7 @@ void mj_transmission(const mjModel* m, mjData* d) {
mji_mulMatVec3(wrench, d->site_xmat+9*refid, gear);
// moment: global Jacobian projected on wrench
mju_mulMatTVec(moment+adr, jac, wrench, 3, nv);
mju_mulMatTVec(moment_row, jac, wrench, 3, nv);
}
// rotational transmission
@@ -1546,18 +1547,18 @@ void mj_transmission(const mjModel* m, mjData* d) {
mjtNum wrench[6];
mji_mulMatVec3(wrench, d->site_xmat+9*refid, gear+3);
// moment_tmp: global Jacobian projected on wrench, add to moment
if (!moment_tmp) moment_tmp = mjSTACKALLOC(d, nv, mjtNum);
mju_mulMatTVec(moment_tmp, jacS, wrench, 3, nv);
mju_addTo(moment+adr, moment_tmp, nv);
// global Jacobian projected on wrench, add to moment
// reuse jac as temporary storage
mju_mulMatTVec(jac, jacS, wrench, 3, nv);
mju_addTo(moment_row, jac, nv);
}
}
// sparsity (compress)
nnz = 0;
for (int j = 0; j < nv; j++) {
if (moment[adr+j]) {
moment[adr+nnz] = moment[adr+j];
if (moment_row[j]) {
moment[adr+nnz] = moment_row[j];
colind[adr+nnz] = j;
nnz++;
}
@@ -1571,7 +1572,8 @@ void mj_transmission(const mjModel* m, mjData* d) {
length[i] = 0;
// clear moment
mju_zero(moment+adr, nv);
if (!moment_row) moment_row = mjSTACKALLOC(d, nv, mjtNum);
mju_zero(moment_row, nv);
// moment is average of all contact normal Jacobians
{
@@ -1655,21 +1657,21 @@ void mj_transmission(const mjModel* m, mjData* d) {
// moment is average over contact normal Jacobians, make negative for adhesion
if (counter) {
// accumulate active contact Jacobians into moment
mj_mulJacTVec(m, d, moment+adr, efc_force);
mj_mulJacTVec(m, d, moment_row, efc_force);
// add Jacobians from excluded contacts
mju_addTo(moment+adr, moment_exclude, nv);
mju_addTo(moment_row, moment_exclude, nv);
// normalize by total contacts, flip sign
mju_scl(moment+adr, moment+adr, -1.0/counter, nv);
mju_scl(moment_row, moment_row, -1.0/counter, nv);
}
}
// sparsity (compress)
nnz = 0;
for (int j = 0; j < nv; j++) {
if (moment[adr+j]) {
moment[adr+nnz] = moment[adr+j];
if (moment_row[j]) {
moment[adr+nnz] = moment_row[j];
colind[adr+nnz] = j;
nnz++;
}
+9 -5
View File
@@ -52,7 +52,7 @@ int mj_isPyramidal(const mjModel* m) {
//-------------------------- sparse chains ---------------------------------------------------------
// merge dof chains for two bodies
int mj_mergeChain(const mjModel* m, int* chain, int b1, int b2) {
int mj_mergeChain(const mjModel* m, int* chain, int b1, int b2, int flg_skipcommon) {
int da1, da2, NV = 0;
// skip fixed bodies
@@ -70,11 +70,15 @@ int mj_mergeChain(const mjModel* m, int* chain, int b1, int b2) {
// merge chains
while (da1 >= 0 || da2 >= 0) {
chain[NV] = mjMAX(da1, da2);
if (da1 == chain[NV]) {
int da = mjMAX(da1, da2);
if (flg_skipcommon && da1 == da && da2 == da) {
break;
}
chain[NV] = da;
if (da1 == da) {
da1 = m->dof_parentid[da1];
}
if (da2 == chain[NV]) {
if (da2 == da) {
da2 = m->dof_parentid[da2];
}
NV++;
@@ -443,7 +447,7 @@ int mj_jacDifPair(const mjModel* m, const mjData* d, int* chain,
if (issimple) {
NV = mj_mergeChainSimple(m, chain, b1, b2);
} else {
NV = mj_mergeChain(m, chain, b1, b2);
NV = mj_mergeChain(m, chain, b1, b2, 0);
}
}
+1 -1
View File
@@ -36,7 +36,7 @@ MJAPI int mj_isSparse(const mjModel* m);
//-------------------------- sparse chains ---------------------------------------------------------
// merge dof chains for two bodies
int mj_mergeChain(const mjModel* m, int* chain, int b1, int b2);
int mj_mergeChain(const mjModel* m, int* chain, int b1, int b2, int flg_skipcommon);
// merge dof chains for two simple bodies
int mj_mergeChainSimple(const mjModel* m, int* chain, int b1, int b2);
+14 -3
View File
@@ -44,6 +44,7 @@
#include <mujoco/mjtnum.h>
#include <mujoco/mujoco.h>
#include "cc/array_safety.h"
#include "engine/engine_core_util.h"
#include "engine/engine_forward.h"
#include "engine/engine_io.h"
#include "engine/engine_name.h"
@@ -3271,9 +3272,13 @@ int mjCModel::CountNJmom(const mjModel* m) {
break;
}
break;
// TODO(taylorhowell): improve upper bounds
case mjTRN_SLIDERCRANK:
count += nv;
{
int id_slider = m->actuator_trnid[2 * i + 1];
std::vector<int> chain(m->nv);
count += mj_mergeChain(m, chain.data(), m->site_bodyid[id],
m->site_bodyid[id_slider], 0);
}
break;
case mjTRN_TENDON:
@@ -3281,7 +3286,13 @@ int mjCModel::CountNJmom(const mjModel* m) {
break;
case mjTRN_SITE:
count += nv;
{
int refid = m->actuator_trnid[2 * i + 1];
int ref_body = refid >= 0 ? m->site_bodyid[refid] : 0;
std::vector<int> chain(m->nv);
count += mj_mergeChain(m, chain.data(), m->site_bodyid[id],
ref_body, /*flg_skipcommon=*/1);
}
break;
case mjTRN_BODY: