Add mj_setKeyframe function to save current state in k-th model keyframe.

Fixes #1719.

Also improve documentation of `mju_sigmoid`, was previously only documented in changelog.

PiperOrigin-RevId: 642637458
Change-Id: I811a0bf8f881c8b76c9a62d2faec8fdec4e7585b
This commit is contained in:
Yuval Tassa
2024-06-12 09:21:54 -07:00
committed by Copybara-Service
parent f37f840880
commit c9bcf8371e
13 changed files with 169 additions and 13 deletions
+18 -1
View File
@@ -235,6 +235,15 @@ mj_setState
Copy concatenated state components specified by ``spec`` from ``state`` into ``d``. The bits of the integer
``spec`` correspond to element fields of :ref:`mjtState`. Fails with :ref:`mju_error` if ``spec`` is invalid.
.. _mj_setKeyframe:
mj_setKeyframe
~~~~~~~~~~~~~~
.. mujoco-include:: mj_setKeyframe
Copy current state to the k-th model keyframe.
.. _mj_addContact:
mj_addContact
@@ -1866,7 +1875,15 @@ mju_sigmoid
.. mujoco-include:: mju_sigmoid
Sigmoid function over 0<=x<=1 using quintic polynomial.
Twice continuously differentiable sigmoid function using a quintic polynomial:
.. math::
s(x) =
\begin{cases}
0, & & x \le 0 \\
6x^5 - 15x^4 + 10x^3, & 0 \lt & x \lt 1 \\
1, & 1 \le & x \qquad
\end{cases}
.. _Interaction:
+12
View File
@@ -533,6 +533,18 @@ Symmetrize square matrix :math:`R = \frac{1}{2}(M + M^T)`.
.. _Miscellaneous:
.. _mju_sigmoid:
Twice continuously differentiable sigmoid function using a quintic polynomial:
.. math::
s(x) =
\begin{cases}
0, & & x \le 0 \\
6x^5 - 15x^4 + 10x^3, & 0 \lt & x \lt 1 \\
1, & 1 \le & x \qquad
\end{cases}
.. _Derivatives-api:
The functions below provide useful derivatives of various functions, both analytic and
+5 -4
View File
@@ -21,13 +21,14 @@ General
:ref:`camera/orthographic<body-camera-orthographic>` and :ref:`global/orthographic<visual-global-orthographic>`
attributes, respectively.
3. Added :ref:`maxhullvert<asset-mesh-maxhullvert>`, the maximum number of vertices in a mesh's convex hull.
4. Add support for ``ball`` joints in the URDF parser.
4. Added :ref:`mj_setKeyframe` for saving the current state into a model keyframe.
5. Added support for ``ball`` joints in the URDF parser ("spherical" in URDF).
MJX
~~~
5. Added support for :ref:`elliptic friction cones<option-cone>`.
6. Fixed a bug that resulted in less-optimal linesearch solutions for some difficult constraint settings.
7. Fixed a bug in the Newton solver that sometimes resulted in less-optimal gradients.
6. Added support for :ref:`elliptic friction cones<option-cone>`.
7. Fixed a bug that resulted in less-optimal linesearch solutions for some difficult constraint settings.
8. Fixed a bug in the Newton solver that sometimes resulted in less-optimal gradients.
Version 3.1.6 (Jun 3, 2024)
---------------------------
+1
View File
@@ -3171,6 +3171,7 @@ void mj_constraintUpdate(const mjModel* m, mjData* d, const mjtNum* jar,
int mj_stateSize(const mjModel* m, unsigned int spec);
void mj_getState(const mjModel* m, const mjData* d, mjtNum* state, unsigned int spec);
void mj_setState(const mjModel* m, mjData* d, const mjtNum* state, unsigned int spec);
void mj_setKeyframe(mjModel* m, const mjData* d, int k);
int mj_addContact(const mjModel* m, mjData* d, const mjContact* con);
int mj_isPyramidal(const mjModel* m);
int mj_isSparse(const mjModel* m);
+3
View File
@@ -418,6 +418,9 @@ MJAPI void mj_getState(const mjModel* m, const mjData* d, mjtNum* state, unsigne
// Set state.
MJAPI void mj_setState(const mjModel* m, mjData* d, const mjtNum* state, unsigned int spec);
// Copy current state to the k-th model keyframe.
MJAPI void mj_setKeyframe(mjModel* m, const mjData* d, int k);
// Add contact to d->contact list; return 0 if success; 1 if buffer full.
MJAPI int mj_addContact(const mjModel* m, mjData* d, const mjContact* con);
+24
View File
@@ -2244,6 +2244,30 @@ FUNCTIONS: Mapping[str, FunctionDecl] = dict([
),
doc='Set state.',
)),
('mj_setKeyframe',
FunctionDecl(
name='mj_setKeyframe',
return_type=ValueType(name='void'),
parameters=(
FunctionParameterDecl(
name='m',
type=PointerType(
inner_type=ValueType(name='mjModel'),
),
),
FunctionParameterDecl(
name='d',
type=PointerType(
inner_type=ValueType(name='mjData', is_const=True),
),
),
FunctionParameterDecl(
name='k',
type=ValueType(name='int'),
),
),
doc='Copy current state to the k-th model keyframe.',
)),
('mj_addContact',
FunctionDecl(
name='mj_addContact',
+32
View File
@@ -27,6 +27,7 @@ import numpy as np
TEST_XML = r"""
<mujoco model="test">
<compiler coordinate="local" angle="radian" eulerseq="xyz"/>
<size nkey="2"/>
<option timestep="0.002" gravity="0 0 -9.81"/>
<visual>
<global fovy="50" />
@@ -745,6 +746,37 @@ class MuJoCoBindingsTest(parameterized.TestCase):
# Expect next states to be equal.
np.testing.assert_array_equal(state1a, state1b)
def test_mj_setKeyframe(self): # pylint: disable=invalid-name
mujoco.mj_step(self.model, self.data)
# Test for invalid state spec
invalid_key = 2
expected_message = (
f'mj_setKeyframe: index must be smaller than {invalid_key} (keyframes'
' allocated in model)'
)
with self.assertRaisesWithLiteralMatch(mujoco.FatalError, expected_message):
mujoco.mj_setKeyframe(self.model, self.data, invalid_key)
valid_key = 1
time = self.data.time
qpos = self.data.qpos.copy()
qvel = self.data.qvel.copy()
act = self.data.act.copy()
mujoco.mj_setKeyframe(self.model, self.data, valid_key)
# Step, assert that time has changed.
mujoco.mj_step(self.model, self.data)
self.assertNotEqual(time, self.data.time)
# Reset to keyframe, assert that time, qpos, qvel, act are the same.
mujoco.mj_resetDataKeyframe(self.model, self.data, valid_key)
self.assertEqual(time, self.data.time)
np.testing.assert_array_equal(qpos, self.data.qpos)
np.testing.assert_array_equal(qvel, self.data.qvel)
np.testing.assert_array_equal(act, self.data.act)
def test_mj_angmomMat(self): # pylint: disable=invalid-name
self.data.qvel = np.ones(self.model.nv, np.float64)
mujoco.mj_forward(self.model, self.data)
+1
View File
@@ -302,6 +302,7 @@ PYBIND11_MODULE(_functions, pymodule) {
}
return InterceptMjErrors(::mj_setState)(m, d, state.data(), spec);
});
Def<traits::mj_setKeyframe>(pymodule);
Def<traits::mj_addContact>(pymodule);
Def<traits::mj_isPyramidal>(pymodule);
Def<traits::mj_isSparse>(pymodule);
+1 -8
View File
@@ -1978,14 +1978,7 @@ void Simulate::Sync() {
}
if (pending_.save_key) {
int i = this->key;
m_->key_time[i] = d_->time;
mju_copy(m_->key_qpos + i*m_->nq, d_->qpos, m_->nq);
mju_copy(m_->key_qvel + i*m_->nv, d_->qvel, m_->nv);
mju_copy(m_->key_act + i*m_->na, d_->act, m_->na);
mju_copy(m_->key_mpos + i*3*m_->nmocap, d_->mocap_pos, 3*m_->nmocap);
mju_copy(m_->key_mquat + i*4*m_->nmocap, d_->mocap_quat, 4*m_->nmocap);
mju_copy(m_->key_ctrl + i*m_->nu, d_->ctrl, m_->nu);
mj_setKeyframe(m_, d_, this->key);
pending_.save_key = false;
}
+22
View File
@@ -232,6 +232,28 @@ void mj_setState(const mjModel* m, mjData* d, const mjtNum* state, unsigned int
// copy current state to the k-th model keyframe
void mj_setKeyframe(mjModel* m, const mjData* d, int k) {
// check keyframe index
if (k >= m->nkey) {
mjERROR("index must be smaller than %d (keyframes allocated in model)", m->nkey);
}
if (k < 0) {
mjERROR("keyframe index cannot be negative");
}
// copy state to model keyframe
m->key_time[k] = d->time;
mju_copy(m->key_qpos + k*m->nq, d->qpos, m->nq);
mju_copy(m->key_qvel + k*m->nv, d->qvel, m->nv);
mju_copy(m->key_act + k*m->na, d->act, m->na);
mju_copy(m->key_mpos + k*3*m->nmocap, d->mocap_pos, 3*m->nmocap);
mju_copy(m->key_mquat + k*4*m->nmocap, d->mocap_quat, 4*m->nmocap);
mju_copy(m->key_ctrl + k*m->nu, d->ctrl, m->nu);
}
//-------------------------- sparse chains ---------------------------------------------------------
// merge dof chains for two bodies
+2
View File
@@ -43,6 +43,8 @@ MJAPI void mj_getState(const mjModel* m, const mjData* d, mjtNum* state, unsigne
// set state
MJAPI void mj_setState(const mjModel* m, mjData* d, const mjtNum* state, unsigned int spec);
// copy current state to the k-th model keyframe
MJAPI void mj_setKeyframe(mjModel* m, const mjData* d, int k);
//-------------------------- sparse chains ---------------------------------------------------------
+45
View File
@@ -722,5 +722,50 @@ TEST_F(SupportTest, GeomDistance) {
mj_deleteModel(model);
}
static constexpr char kSetKeyframeTestingModel[] = R"(
<mujoco>
<size nkey="2"/>
<worldbody>
<body>
<joint name="joint" axis="0 1 0"/>
<geom size=".1" pos="1 0 0"/>
</body>
</worldbody>
<actuator>
<intvelocity joint="joint" actrange="-1 1" kp="100" dampratio="1"/>
</actuator>
</mujoco>
)";
TEST_F(SupportTest, SetKeyframe) {
mjModel* model = LoadModelFromString(kSetKeyframeTestingModel);
mjData* data = mj_makeData(model);
data->ctrl[0] = 1;
while (data->time < 1) {
mj_step(model, data);
}
mj_setKeyframe(model, data, 1);
EXPECT_EQ(data->time, model->key_time[1]);
EXPECT_EQ(data->ctrl[0], model->key_ctrl[model->nu * 1]);
EXPECT_EQ(data->qpos[0], model->key_qpos[model->nq * 1]);
EXPECT_EQ(data->qvel[0], model->key_qvel[model->nv * 1]);
EXPECT_EQ(data->act[0], model->key_act[model->na * 1]);
mj_step(model, data);
mj_setKeyframe(model, data, 0);
EXPECT_EQ(data->time, model->key_time[0]);
EXPECT_EQ(data->ctrl[0], model->key_ctrl[model->nu * 0]);
EXPECT_EQ(data->qpos[0], model->key_qpos[model->nq * 0]);
EXPECT_EQ(data->qvel[0], model->key_qvel[model->nv * 0]);
EXPECT_EQ(data->act[0], model->key_act[model->na * 0]);
mj_deleteData(data);
mj_deleteModel(model);
}
} // namespace
} // namespace mujoco
+3
View File
@@ -6641,6 +6641,9 @@ public static unsafe extern void mj_getState(mjModel_* m, mjData_* d, double* st
[DllImport("mujoco", CallingConvention = CallingConvention.Cdecl)]
public static unsafe extern void mj_setState(mjModel_* m, mjData_* d, double* state, uint spec);
[DllImport("mujoco", CallingConvention = CallingConvention.Cdecl)]
public static unsafe extern void mj_setKeyframe(mjModel_* m, mjData_* d, int k);
[DllImport("mujoco", CallingConvention = CallingConvention.Cdecl)]
public static unsafe extern int mj_addContact(mjModel_* m, mjData_* d, mjContact_* con);