From ec43ec7c301b0b4b485ec9e2fb89300250dd3737 Mon Sep 17 00:00:00 2001 From: Alessio Quaglino Date: Fri, 9 Aug 2024 12:52:48 -0700 Subject: [PATCH] Remove hard-coded nPOS and nVEL from mjCModel. PiperOrigin-RevId: 661370196 Change-Id: Ib2c1876dfea55dbf496e2358000f24bbde29fc52 --- src/user/user_model.cc | 51 +++++++++------------------------------- src/user/user_objects.cc | 28 ++++++++++++++++++++++ src/user/user_objects.h | 5 ++++ 3 files changed, 44 insertions(+), 40 deletions(-) diff --git a/src/user/user_model.cc b/src/user/user_model.cc index c7fd5ccb..3173da06 100644 --- a/src/user/user_model.cc +++ b/src/user/user_model.cc @@ -1360,10 +1360,6 @@ void mjCModel::CheckEmptyNames(void) { -// number of position and velocity coordinates for each joint type -const int nPOS[4] = {7, 4, 1, 1}; -const int nVEL[4] = {6, 3, 1, 1}; - template static size_t getpathslength(std::vector list) { size_t result = 0; @@ -1404,8 +1400,8 @@ void mjCModel::SetSizes() { // nq, nv for (int i=0; itype]; - nv += nVEL[joints_[i]->type]; + nq += joints_[i]->nq(); + nv += joints_[i]->nv(); } // nu, na @@ -1536,7 +1532,7 @@ void mjCModel::AutoSpringDamper(mjModel* m) { for (int n=0; nnjnt; n++) { // get joint dof address and number of dimensions int adr = m->jnt_dofadr[n]; - int ndim = nVEL[m->jnt_type[n]]; + int ndim = mjCJoint::nv((mjtJoint)m->jnt_type[n]); // get timeconst and dampratio from joint specificatin mjtNum timeconst = (mjtNum)joints_[n]->springdamper[0]; @@ -2009,7 +2005,7 @@ void mjCModel::CopyTree(mjModel* m) { } // set dof fields for this joint - for (int j1=0; j1type]; j1++) { + for (int j1=0; j1nv(); j1++) { // set attributes m->dof_bodyid[dofadr] = pb->id; m->dof_jntid[dofadr] = jid; @@ -2029,7 +2025,7 @@ void mjCModel::CopyTree(mjModel* m) { // advance joint and qpos counters jntadr++; - qposadr += nPOS[pj->type]; + qposadr += pj->nq(); } // simple body with sliders and no rotational dofs: promote to simple level 2 @@ -2832,21 +2828,8 @@ void mjCModel::CopyObjects(mjModel* m) { // save the current state void mjCModel::SaveState(const mjtNum* qpos, const mjtNum* qvel, const mjtNum* act) { for (auto joint : joints_) { - switch (joint->type) { - case mjJNT_FREE: - if (qpos) mjuu_copyvec(joint->qpos, qpos + joint->qposadr_, 7); - if (qvel) mjuu_copyvec(joint->qvel, qvel + joint->dofadr_, 6); - break; - case mjJNT_BALL: - if (qpos) mjuu_copyvec(joint->qpos, qpos + joint->qposadr_, 4); - if (qvel) mjuu_copyvec(joint->qvel, qvel + joint->dofadr_, 3); - break; - case mjJNT_HINGE: - case mjJNT_SLIDE: - if (qpos) mjuu_copyvec(joint->qpos, qpos + joint->qposadr_, 1); - if (qvel) mjuu_copyvec(joint->qvel, qvel + joint->dofadr_, 1); - break; - } + if (qpos) mjuu_copyvec(joint->qpos, qpos + joint->qposadr_, joint->nq()); + if (qvel) mjuu_copyvec(joint->qvel, qvel + joint->dofadr_, joint->nv()); } for (auto actuator : actuators_) { @@ -2874,23 +2857,11 @@ void mjCModel::MakeData(const mjModel* m, mjData** dest) { // restore the previous state void mjCModel::RestoreState(mjtNum* qpos, mjtNum* qvel, mjtNum* act) { for (auto joint : joints_) { - if (!mjuu_defined(joint->qpos[0]) || !mjuu_defined(joint->qvel[0])) { - continue; + if (mjuu_defined(joint->qpos[0]) && qpos) { + mjuu_copyvec(qpos + joint->qposadr_, joint->qpos, joint->nq()); } - switch (joint->type) { - case mjJNT_FREE: - if (qpos) mjuu_copyvec(qpos + joint->qposadr_, joint->qpos, 7); - if (qvel) mjuu_copyvec(qvel + joint->dofadr_, joint->qvel, 6); - break; - case mjJNT_BALL: - if (qpos) mjuu_copyvec(qpos + joint->qposadr_, joint->qpos, 4); - if (qvel) mjuu_copyvec(qvel + joint->dofadr_, joint->qvel, 3); - break; - case mjJNT_HINGE: - case mjJNT_SLIDE: - if (qpos) mjuu_copyvec(qpos + joint->qposadr_, joint->qpos, 1); - if (qvel) mjuu_copyvec(qvel + joint->dofadr_, joint->qvel, 1); - break; + if (mjuu_defined(joint->qvel[0]) && qvel) { + mjuu_copyvec(qvel + joint->dofadr_, joint->qvel, joint->nv()); } } diff --git a/src/user/user_objects.cc b/src/user/user_objects.cc index 0bdc4f09..91322954 100644 --- a/src/user/user_objects.cc +++ b/src/user/user_objects.cc @@ -1778,6 +1778,34 @@ bool mjCJoint::is_actfrclimited() const { return islimited(actfrclimited, actfrc +int mjCJoint::nq(mjtJoint joint_type) { + switch (joint_type) { + case mjJNT_FREE: + return 7; + case mjJNT_BALL: + return 4; + case mjJNT_SLIDE: + case mjJNT_HINGE: + return 1; + } +} + + + +int mjCJoint::nv(mjtJoint joint_type) { + switch (joint_type) { + case mjJNT_FREE: + return 6; + case mjJNT_BALL: + return 3; + case mjJNT_SLIDE: + case mjJNT_HINGE: + return 1; + } +} + + + void mjCJoint::PointToLocal() { spec.element = static_cast(this); spec.name = &name; diff --git a/src/user/user_objects.h b/src/user/user_objects.h index 18d9b886..a6540a81 100644 --- a/src/user/user_objects.h +++ b/src/user/user_objects.h @@ -424,6 +424,11 @@ class mjCJoint : public mjCJoint_, private mjsJoint { bool is_limited() const; bool is_actfrclimited() const; + static int nq(mjtJoint joint_type); + static int nv(mjtJoint joint_type); + int nq() const { return nq(spec.type); } + int nv() const { return nv(spec.type); } + private: int Compile(void); // compiler; return dofnum void PointToLocal(void);