Allow passing a null pointer in mj_recompile.

Fixes #1819.

PiperOrigin-RevId: 653242587
Change-Id: I36dbd6ed54abdf21abeb62cf41ee88f159008d01
This commit is contained in:
Alessio Quaglino
2024-07-17 08:28:13 -07:00
committed by Copybara-Service
parent a2fd34caee
commit 2664cc1391
4 changed files with 38 additions and 24 deletions
+7 -2
View File
@@ -86,9 +86,14 @@ mjModel* mj_compile(mjSpec* s, const mjVFS* vfs) {
// recompile spec into existing model and data while preserving the state
void mj_recompile(mjSpec* s, const mjVFS* vfs, mjModel* m, mjData* d) {
mjCModel* modelC = static_cast<mjCModel*>(s->element);
modelC->SaveState(d);
if (d) {
modelC->SaveState(d->qpos, d->qvel, d->act);
}
modelC->Compile(vfs, &m);
modelC->RestoreState(m, &d);
if (d) {
modelC->MakeData(m, &d);
modelC->RestoreState(d->qpos, d->qvel, d->act);
}
}
+24 -19
View File
@@ -2776,68 +2776,73 @@ void mjCModel::CopyObjects(mjModel* m) {
// save the current state
void mjCModel::SaveState(const mjData* d) {
void mjCModel::SaveState(const mjtNum* qpos, const mjtNum* qvel, const mjtNum* act) {
for (auto joint : joints_) {
switch (joint->type) {
case mjJNT_FREE:
mjuu_copyvec(joint->qpos, d->qpos + joint->qposadr_, 7);
mjuu_copyvec(joint->qvel, d->qvel + joint->dofadr_, 6);
if (qpos) mjuu_copyvec(joint->qpos, qpos + joint->qposadr_, 7);
if (qvel) mjuu_copyvec(joint->qvel, qvel + joint->dofadr_, 6);
break;
case mjJNT_BALL:
mjuu_copyvec(joint->qpos, d->qpos + joint->qposadr_, 4);
mjuu_copyvec(joint->qvel, d->qvel + joint->dofadr_, 3);
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:
mjuu_copyvec(joint->qpos, d->qpos + joint->qposadr_, 1);
mjuu_copyvec(joint->qvel, d->qvel + joint->dofadr_, 1);
if (qpos) mjuu_copyvec(joint->qpos, qpos + joint->qposadr_, 1);
if (qvel) mjuu_copyvec(joint->qvel, qvel + joint->dofadr_, 1);
break;
}
}
for (auto actuator : actuators_) {
if (actuator->actadr_ != -1) {
if (actuator->actadr_ != -1 && act) {
actuator->act.assign(actuator->actnum_, 0);
mjuu_copyvec(actuator->act.data(), d->act + actuator->actadr_, actuator->actnum_);
mjuu_copyvec(actuator->act.data(), act + actuator->actadr_, actuator->actnum_);
}
}
}
// restore the previous state
void mjCModel::RestoreState(const mjModel* m, mjData** dest) {
// clear existing data
void mjCModel::MakeData(const mjModel* m, mjData** dest) {
mj_makeRawData(dest, m);
mjData* d = *dest;
if (d) {
mj_initPlugin(m, d);
mj_resetData(m, d);
}
}
// 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;
}
switch (joint->type) {
case mjJNT_FREE:
mjuu_copyvec(d->qpos + joint->qposadr_, joint->qpos, 7);
mjuu_copyvec(d->qvel + joint->dofadr_, joint->qvel, 6);
if (qpos) mjuu_copyvec(qpos + joint->qposadr_, joint->qpos, 7);
if (qvel) mjuu_copyvec(qvel + joint->dofadr_, joint->qvel, 6);
break;
case mjJNT_BALL:
mjuu_copyvec(d->qpos + joint->qposadr_, joint->qpos, 4);
mjuu_copyvec(d->qvel + joint->dofadr_, joint->qvel, 3);
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:
mjuu_copyvec(d->qpos + joint->qposadr_, joint->qpos, 1);
mjuu_copyvec(d->qvel + joint->dofadr_, joint->qvel, 1);
if (qpos) mjuu_copyvec(qpos + joint->qposadr_, joint->qpos, 1);
if (qvel) mjuu_copyvec(qvel + joint->dofadr_, joint->qvel, 1);
break;
}
}
for (auto actuator : actuators_) {
if (mjuu_defined(actuator->act[0])) {
mjuu_copyvec(d->act + actuator->actadr_, actuator->act.data(), actuator->actnum_);
if (mjuu_defined(actuator->act[0]) && act) {
mjuu_copyvec(act + actuator->actadr_, actuator->act.data(), actuator->actnum_);
}
}
}
+3 -2
View File
@@ -273,8 +273,9 @@ class mjCModel : public mjCModel_, private mjSpec {
std::string_view name = "");
// save/restore the current state
void SaveState(const mjData* d);
void RestoreState(const mjModel* m, mjData** dest);
void SaveState(const mjtNum* qpos, const mjtNum* qvel, const mjtNum* act);
void MakeData(const mjModel* m, mjData** dest);
void RestoreState(mjtNum* qpos, mjtNum* qvel, mjtNum* act);
// map from default class name to default class pointer
std::unordered_map<std::string, mjCDef*> def_map;
+4 -1
View File
@@ -863,8 +863,11 @@ TEST_F(MujocoTest, PreserveState) {
EXPECT_EQ(data->act[i], d_expected->act[i]) << i;
}
// destroy everything
// check that the function is callable with no data
mj_deleteData(data);
mj_recompile(spec, 0, model, nullptr);
// destroy everything
mj_deleteData(d_expected);
mj_deleteSpec(spec);
mj_deleteModel(model);