Restore full state in mjd_stepFD.
Make sure that plugin state is also restored after finite-difference calculations. PiperOrigin-RevId: 573272654 Change-Id: Ie276ecc391c463ae5c47cafa5b8de873b9345a7d
This commit is contained in:
committed by
Copybara-Service
parent
fb4cf472ed
commit
eb41c720d4
@@ -35,11 +35,7 @@
|
||||
|
||||
// get state=[qpos; qvel; act] and optionally sensordata
|
||||
static void getState(const mjModel* m, const mjData* d, mjtNum* state, mjtNum* sensordata) {
|
||||
int nq = m->nq, nv = m->nv, na = m->na;
|
||||
|
||||
mju_copy(state, d->qpos, nq);
|
||||
mju_copy(state+nq, d->qvel, nv);
|
||||
mju_copy(state+nq+nv, d->act, na);
|
||||
mj_getState(m, d, state, mjSTATE_PHYSICS);
|
||||
if (sensordata) {
|
||||
mju_copy(sensordata, d->sensordata, m->nsensordata);
|
||||
}
|
||||
@@ -47,25 +43,6 @@ static void getState(const mjModel* m, const mjData* d, mjtNum* state, mjtNum* s
|
||||
|
||||
|
||||
|
||||
// set state=[qpos; qvel; act] and optionally warmstart accelerations
|
||||
static void setState(const mjModel* m, mjData* d, mjtNum time, const mjtNum* state,
|
||||
const mjtNum* ctrl, const mjtNum* warmstart) {
|
||||
int nq = m->nq, nv = m->nv, na = m->na;
|
||||
|
||||
d->time = time;
|
||||
mju_copy(d->qpos, state, nq);
|
||||
mju_copy(d->qvel, state+nq, nv);
|
||||
mju_copy(d->act, state+nq+nv, na);
|
||||
if (ctrl) {
|
||||
mju_copy(d->ctrl, ctrl, m->nu);
|
||||
}
|
||||
if (warmstart) {
|
||||
mju_copy(d->qacc_warmstart, warmstart, nv);
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
|
||||
// dx = (x2 - x1) / h
|
||||
static void diff(mjtNum* restrict dx, const mjtNum* x1, const mjtNum* x2, mjtNum h, int n) {
|
||||
mjtNum inv_h = 1/h;
|
||||
@@ -333,14 +310,17 @@ void mjd_stepFD(const mjModel* m, mjData* d, mjtNum eps, mjtByte flg_centered,
|
||||
mj_markStack(d);
|
||||
|
||||
// states
|
||||
// state to restore after finite differencing
|
||||
unsigned int restore_spec = mjSTATE_FULLPHYSICS | mjSTATE_CTRL;
|
||||
|
||||
restore_spec |= mjDISABLED(mjDSBL_WARMSTART) ? 0 : mjSTATE_WARMSTART;
|
||||
|
||||
mjtNum *fullstate = mj_stackAllocNum(d, mj_stateSize(m, restore_spec));
|
||||
mjtNum *state = mj_stackAllocNum(d, nq+nv+na); // current state
|
||||
mjtNum *next = mj_stackAllocNum(d, nq+nv+na); // next state
|
||||
mjtNum *next_plus = mj_stackAllocNum(d, nq+nv+na); // forward-nudged next state
|
||||
mjtNum *next_minus = mj_stackAllocNum(d, nq+nv+na); // backward-nudged next state
|
||||
|
||||
// warmstart accelerations
|
||||
mjtNum *warmstart = mjDISABLED(mjDSBL_WARMSTART) ? NULL : mj_stackAllocNum(d, nv);
|
||||
|
||||
// sensors
|
||||
int skipsensor = !DsDq && !DsDv && !DsDa && !DsDu;
|
||||
mjtNum *sensor = skipsensor ? NULL : mj_stackAllocNum(d, ns); // sensor values
|
||||
@@ -351,12 +331,9 @@ void mjd_stepFD(const mjModel* m, mjData* d, mjtNum eps, mjtByte flg_centered,
|
||||
mjtNum *ctrl = mj_stackAllocNum(d, nu);
|
||||
|
||||
// save current inputs
|
||||
mjtNum time = d->time;
|
||||
mj_getState(m, d, fullstate, restore_spec);
|
||||
mju_copy(ctrl, d->ctrl, nu);
|
||||
getState(m, d, state, NULL);
|
||||
if (warmstart) {
|
||||
mju_copy(warmstart, d->qacc_warmstart, nv);
|
||||
}
|
||||
|
||||
// step input
|
||||
mj_stepSkip(m, d, mjSTAGE_NONE, skipsensor);
|
||||
@@ -365,7 +342,7 @@ void mjd_stepFD(const mjModel* m, mjData* d, mjtNum eps, mjtByte flg_centered,
|
||||
getState(m, d, next, sensor);
|
||||
|
||||
// restore input
|
||||
setState(m, d, time, state, ctrl, warmstart);
|
||||
mj_setState(m, d, fullstate, restore_spec);
|
||||
|
||||
// finite-difference controls: skip=mjSTAGE_VEL, handle ctrl at range limits
|
||||
if (DyDu || DsDu) {
|
||||
@@ -382,7 +359,7 @@ void mjd_stepFD(const mjModel* m, mjData* d, mjtNum eps, mjtByte flg_centered,
|
||||
getState(m, d, next_plus, sensor_plus);
|
||||
|
||||
// reset
|
||||
setState(m, d, time, state, ctrl, warmstart);
|
||||
mj_setState(m, d, fullstate, restore_spec);
|
||||
}
|
||||
|
||||
// nudge backward, if possible given ctrlrange
|
||||
@@ -397,7 +374,7 @@ void mjd_stepFD(const mjModel* m, mjData* d, mjtNum eps, mjtByte flg_centered,
|
||||
getState(m, d, next_minus, sensor_minus);
|
||||
|
||||
// reset
|
||||
setState(m, d, time, state, ctrl, warmstart);
|
||||
mj_setState(m, d, fullstate, restore_spec);
|
||||
}
|
||||
|
||||
// difference states
|
||||
@@ -425,7 +402,7 @@ void mjd_stepFD(const mjModel* m, mjData* d, mjtNum eps, mjtByte flg_centered,
|
||||
getState(m, d, next_plus, sensor_plus);
|
||||
|
||||
// reset
|
||||
setState(m, d, time, state, NULL, warmstart);
|
||||
mj_setState(m, d, fullstate, restore_spec);
|
||||
|
||||
// nudge backward
|
||||
if (flg_centered) {
|
||||
@@ -437,7 +414,7 @@ void mjd_stepFD(const mjModel* m, mjData* d, mjtNum eps, mjtByte flg_centered,
|
||||
getState(m, d, next_minus, sensor_minus);
|
||||
|
||||
// reset
|
||||
setState(m, d, time, state, NULL, warmstart);
|
||||
mj_setState(m, d, fullstate, restore_spec);
|
||||
}
|
||||
|
||||
// difference states
|
||||
@@ -472,7 +449,7 @@ void mjd_stepFD(const mjModel* m, mjData* d, mjtNum eps, mjtByte flg_centered,
|
||||
getState(m, d, next_plus, sensor_plus);
|
||||
|
||||
// reset
|
||||
setState(m, d, time, state, NULL, warmstart);
|
||||
mj_setState(m, d, fullstate, restore_spec);
|
||||
|
||||
// nudge backward
|
||||
if (flg_centered) {
|
||||
@@ -484,7 +461,7 @@ void mjd_stepFD(const mjModel* m, mjData* d, mjtNum eps, mjtByte flg_centered,
|
||||
getState(m, d, next_minus, sensor_minus);
|
||||
|
||||
// reset
|
||||
setState(m, d, time, state, NULL, warmstart);
|
||||
mj_setState(m, d, fullstate, restore_spec);
|
||||
}
|
||||
|
||||
// difference states
|
||||
@@ -521,7 +498,7 @@ void mjd_stepFD(const mjModel* m, mjData* d, mjtNum eps, mjtByte flg_centered,
|
||||
getState(m, d, next_plus, sensor_plus);
|
||||
|
||||
// reset
|
||||
setState(m, d, time, state, NULL, warmstart);
|
||||
mj_setState(m, d, fullstate, restore_spec);
|
||||
|
||||
// nudge backward
|
||||
if (flg_centered) {
|
||||
@@ -535,7 +512,7 @@ void mjd_stepFD(const mjModel* m, mjData* d, mjtNum eps, mjtByte flg_centered,
|
||||
getState(m, d, next_minus, sensor_minus);
|
||||
|
||||
// reset
|
||||
setState(m, d, time, state, NULL, warmstart);
|
||||
mj_setState(m, d, fullstate, restore_spec);
|
||||
}
|
||||
|
||||
// difference states
|
||||
|
||||
Reference in New Issue
Block a user