diff --git a/src/user/user_api.cc b/src/user/user_api.cc index 1a03b274..7ff36d73 100644 --- a/src/user/user_api.cc +++ b/src/user/user_api.cc @@ -96,12 +96,12 @@ void mj_recompile(mjSpec* s, const mjVFS* vfs, mjModel* m, mjData* d) { mjtNum time = 0; if (d) { time = d->time; - modelC->SaveState(d->qpos, d->qvel, d->act); + modelC->SaveState(d->qpos, d->qvel, d->act, d->ctrl); } modelC->Compile(vfs, &m); if (d) { modelC->MakeData(m, &d); - modelC->RestoreState(m->qpos0, d->qpos, d->qvel, d->act); + modelC->RestoreState(m->qpos0, d->qpos, d->qvel, d->act, d->ctrl); d->time = time; } } diff --git a/src/user/user_model.cc b/src/user/user_model.cc index 5ea2124c..b96189be 100644 --- a/src/user/user_model.cc +++ b/src/user/user_model.cc @@ -2863,7 +2863,7 @@ void mjCModel::CopyObjects(mjModel* m) { // save the current state template -void mjCModel::SaveState(const T* qpos, const T* qvel, const T* act) { +void mjCModel::SaveState(const T* qpos, const T* qvel, const T* act, const T* ctrl) { for (auto joint : joints_) { if (joint->qposadr_ == -1 || joint->dofadr_ == -1) { throw mjCError(NULL, "SaveState: joint %s has no address", joint->name.c_str()); @@ -2872,11 +2872,15 @@ void mjCModel::SaveState(const T* qpos, const T* qvel, const T* act) { if (qvel) mjuu_copyvec(joint->qvel(), qvel + joint->dofadr_, joint->nv()); } - for (auto actuator : actuators_) { + for (unsigned int i=0; iactadr_ != -1 && actuator->actdim_ != -1 && act) { actuator->act().assign(actuator->actdim_, 0); mjuu_copyvec(actuator->act().data(), act + actuator->actadr_, actuator->actdim_); } + if (ctrl) { + actuator->ctrl() = ctrl[i]; + } } } @@ -2896,7 +2900,7 @@ void mjCModel::MakeData(const mjModel* m, mjData** dest) { // restore the previous state template -void mjCModel::RestoreState(const mjtNum* pos0, T* qpos, T* qvel, T* act) { +void mjCModel::RestoreState(const mjtNum* pos0, T* qpos, T* qvel, T* act, T* ctrl) { for (auto joint : joints_) { if (qpos) { if (mjuu_defined(joint->qpos()[0])) { @@ -2911,10 +2915,14 @@ void mjCModel::RestoreState(const mjtNum* pos0, T* qpos, T* qvel, T* act) { } // restore act - for (auto actuator : actuators_) { + for (unsigned int i=0; iact().empty() && mjuu_defined(actuator->act()[0]) && act) { mjuu_copyvec(act + actuator->actadr_, actuator->act().data(), actuator->actdim_); } + if (ctrl) { + ctrl[i] = mjuu_defined(actuator->ctrl()) ? actuator->ctrl() : 0; + } } } @@ -2923,10 +2931,11 @@ void mjCModel::RestoreState(const mjtNum* pos0, T* qpos, T* qvel, T* act) { // force explicit instantiations template void mjCModel::SaveState(const mjtNum* qpos, const mjtNum* qvel, - const mjtNum* act); + const mjtNum* act, + const mjtNum* ctrl); template void mjCModel::RestoreState(const mjtNum* qpos0, mjtNum* qpos, - mjtNum* qvel, mjtNum* act); + mjtNum* qvel, mjtNum* act, mjtNum* ctrl); @@ -2946,9 +2955,11 @@ void mjCModel::StoreKeyframes() { info.qpos = !key->spec_qpos_.empty(); info.qvel = !key->spec_qvel_.empty(); info.act = !key->spec_act_.empty(); + info.ctrl = !key->spec_ctrl_.empty(); key_pending_.push_back(info); state_name_ = info.name; - SaveState(key->spec_qpos_.data(), key->spec_qvel_.data(), key->spec_act_.data()); + SaveState(key->spec_qpos_.data(), key->spec_qvel_.data(), + key->spec_act_.data(), key->spec_ctrl_.data()); } if (resetlists) { @@ -3506,6 +3517,9 @@ void mjCModel::ResolveKeyframes(const mjModel* m) { if (!key->spec_act_.empty()) { key->spec_act_.resize(na); } + if (!key->spec_ctrl_.empty()) { + key->spec_ctrl_.resize(nu); + } } // store dof offsets in joints and actuators @@ -3518,8 +3532,10 @@ void mjCModel::ResolveKeyframes(const mjModel* m) { if (info.qpos) key->spec_qpos_.assign(nq, 0); if (info.qvel) key->spec_qvel_.assign(nv, 0); if (info.act) key->spec_act_.assign(na, 0); + if (info.ctrl) key->spec_ctrl_.assign(nu, 0); state_name_ = info.name; - RestoreState(m->qpos0, key->spec_qpos_.data(), key->spec_qvel_.data(), key->spec_act_.data()); + RestoreState(m->qpos0, key->spec_qpos_.data(), key->spec_qvel_.data(), + key->spec_act_.data(), key->spec_ctrl_.data()); } // the attached keyframes have been copied into the model diff --git a/src/user/user_model.h b/src/user/user_model.h index 674ab72a..8b004b79 100644 --- a/src/user/user_model.h +++ b/src/user/user_model.h @@ -39,6 +39,7 @@ typedef struct mjKeyInfo_ { bool qpos; bool qvel; bool act; + bool ctrl; } mjKeyInfo; class mjCModel_ : public mjsElement { @@ -283,8 +284,8 @@ class mjCModel : public mjCModel_, private mjSpec { std::string_view name = ""); // save/restore the current state - template void SaveState(const T* qpos, const T* qvel, const T* act); - template void RestoreState(const mjtNum* pos0, T* qpos, T* qvel, T* act); + template void SaveState(const T* qpos, const T* qvel, const T* act, const T* ctrl); + template void RestoreState(const mjtNum* pos0, T* qpos, T* qvel, T* act, T* ctrl); // clear existing data void MakeData(const mjModel* m, mjData** dest); diff --git a/src/user/user_objects.cc b/src/user/user_objects.cc index 0afb39e7..12f55c1f 100644 --- a/src/user/user_objects.cc +++ b/src/user/user_objects.cc @@ -5394,6 +5394,7 @@ mjCActuator& mjCActuator::operator=(const mjCActuator& other) { void mjCActuator::ForgetKeyframes() { act_.clear(); + ctrl_.clear(); } @@ -5413,6 +5414,15 @@ std::vector& mjCActuator::act() { +mjtNum& mjCActuator::ctrl() { + if (ctrl_.find(model->state_name_) == ctrl_.end()) { + ctrl_[model->state_name_] = mjNAN; + } + return ctrl_.at(model->state_name_); +} + + + void mjCActuator::PointToLocal() { spec.element = static_cast(this); spec.name = &name; diff --git a/src/user/user_objects.h b/src/user/user_objects.h index bd39c7ef..62401489 100644 --- a/src/user/user_objects.h +++ b/src/user/user_objects.h @@ -1397,6 +1397,7 @@ class mjCActuator_ : public mjCBase { int actadr_; // address of dof in data->act int actdim_; // number of dofs in data->act std::map> act_; // act at the previous step + std::map ctrl_; // ctrl at the previous step // variable-size data std::string plugin_name; @@ -1436,6 +1437,7 @@ class mjCActuator : public mjCActuator_, private mjsActuator { bool is_actlimited() const; std::vector& act(); + mjtNum& ctrl(); private: void Compile(void); // compiler diff --git a/test/user/user_api_test.cc b/test/user/user_api_test.cc index 75280243..9024be04 100644 --- a/test/user/user_api_test.cc +++ b/test/user/user_api_test.cc @@ -430,8 +430,8 @@ static constexpr char xml_child[] = R"( - - + + )"; @@ -501,10 +501,10 @@ TEST_F(MujocoTest, AttachSame) { - - - - + + + + )"; @@ -621,8 +621,8 @@ TEST_F(MujocoTest, AttachDifferent) { - - + + )"; @@ -743,8 +743,8 @@ TEST_F(MujocoTest, AttachFrame) { - - + + )";