Add mj_extractState as a public API function.

This function allows the caller to extract a subset of components of a state previously obtained via `mj_getState` without having to first write it back into `mjData`.

PiperOrigin-RevId: 822165077
Change-Id: I6433261f4f5ec2e8d024bb7dee8f7dbf1e62666c
This commit is contained in:
Saran Tunyasuvunakool
2025-10-21 10:00:24 -07:00
committed by Copybara-Service
parent 25f5be82ff
commit 2f65e23779
11 changed files with 171 additions and 1 deletions
+11
View File
@@ -259,6 +259,17 @@ correspond to element fields of :ref:`mjtState`.
Copy concatenated state components specified by ``sig`` from ``d`` into ``state``. The bits of the integer
``sig`` correspond to element fields of :ref:`mjtState`. Fails with :ref:`mju_error` if ``sig`` is invalid.
.. _mj_extractState:
`mj_extractState <#mj_extractState>`__
~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~
.. mujoco-include:: mj_extractState
Extract into ``dst`` the subset of components specified by ``dstsig`` from a state ``src`` previously obtained via
:ref:`mj_getState` with components specified by ``srcsig``. Fails with :ref:`mju_error` if the bits set in ``dstsig``
is not a subset of the bits set in ``srcsig``.
.. _mj_setState:
`mj_setState <#mj_setState>`__
+6
View File
@@ -195,6 +195,12 @@ correspond to element fields of :ref:`mjtState`.
Copy concatenated state components specified by ``sig`` from ``d`` into ``state``. The bits of the integer
``sig`` correspond to element fields of :ref:`mjtState`. Fails with :ref:`mju_error` if ``sig`` is invalid.
.. _mj_extractState:
Extract into ``dst`` the subset of components specified by ``dstsig`` from a state ``src`` previously obtained via
:ref:`mj_getState` with components specified by ``srcsig``. Fails with :ref:`mju_error` if the bits set in ``dstsig``
is not a subset of the bits set in ``srcsig``.
.. _mj_setState:
Copy concatenated state components specified by ``sig`` from ``state`` into ``d``. The bits of the integer
+2
View File
@@ -10,6 +10,8 @@ General
- Raise an error if there are name collisions also during parsing.
- Increase Windows stack size to 16MB to enable models with deep nested body hierarchies.
- Added a new :ref:`mj_extractState` function that allows a subset of a state that was previously returned by
:ref:`mj_getState` to be extracted without having to be written back into ``mjData`` first.
Version 3.3.7 (October 13, 2025)
-----------------------------------
+2
View File
@@ -3118,6 +3118,8 @@ void mj_constraintUpdate(const mjModel* m, mjData* d, const mjtNum* jar,
mjtNum cost[1], int flg_coneHessian);
int mj_stateSize(const mjModel* m, unsigned int sig);
void mj_getState(const mjModel* m, const mjData* d, mjtNum* state, unsigned int sig);
void mj_extractState(const mjModel* m, const mjtNum* src, unsigned int srcsig,
mjtNum* dst, unsigned int dstsig);
void mj_setState(const mjModel* m, mjData* d, const mjtNum* state, unsigned int sig);
void mj_setKeyframe(mjModel* m, const mjData* d, int k);
int mj_addContact(const mjModel* m, mjData* d, const mjContact* con);
+4
View File
@@ -465,6 +465,10 @@ MJAPI int mj_stateSize(const mjModel* m, unsigned int sig);
// Get state.
MJAPI void mj_getState(const mjModel* m, const mjData* d, mjtNum* state, unsigned int sig);
// Extract a subset of components from a state previously obtained via mj_getState.
MJAPI void mj_extractState(const mjModel* m, const mjtNum* src, unsigned int srcsig,
mjtNum* dst, unsigned int dstsig);
// Set state.
MJAPI void mj_setState(const mjModel* m, mjData* d, const mjtNum* state, unsigned int sig);
+15 -1
View File
@@ -329,10 +329,24 @@ PYBIND11_MODULE(_functions, pymodule) {
}
return InterceptMjErrors(::mj_getState)(m, d, state.data(), sig);
});
Def<traits::mj_extractState>(
pymodule,
[](const raw::MjModel* m,
Eigen::Ref<const EigenVectorX> src, unsigned int srcsig,
Eigen::Ref<EigenVectorX> dst, unsigned int dstsig) {
if (src.size() != mj_stateSize(m, srcsig)) {
throw py::type_error("src size should equal mj_stateSize(m, srcsig)");
}
if (dst.size() != mj_stateSize(m, dstsig)) {
throw py::type_error("dst size should equal mj_stateSize(m, dstsig)");
}
return InterceptMjErrors(::mj_extractState)(m, src.data(), srcsig,
dst.data(), dstsig);
});
Def<traits::mj_setState>(
pymodule,
[](const raw::MjModel* m, raw::MjData* d,
const Eigen::Ref<EigenVectorX> state, unsigned int sig) {
Eigen::Ref<const EigenVectorX> state, unsigned int sig) {
if (state.size() != mj_stateSize(m, sig)) {
throw py::type_error("state size should equal mj_stateSize(m, sig)");
}
+34
View File
@@ -2416,6 +2416,40 @@ FUNCTIONS: Mapping[str, FunctionDecl] = dict([
),
doc='Get state.',
)),
('mj_extractState',
FunctionDecl(
name='mj_extractState',
return_type=ValueType(name='void'),
parameters=(
FunctionParameterDecl(
name='m',
type=PointerType(
inner_type=ValueType(name='mjModel', is_const=True),
),
),
FunctionParameterDecl(
name='src',
type=PointerType(
inner_type=ValueType(name='mjtNum', is_const=True),
),
),
FunctionParameterDecl(
name='srcsig',
type=ValueType(name='unsigned int'),
),
FunctionParameterDecl(
name='dst',
type=PointerType(
inner_type=ValueType(name='mjtNum'),
),
),
FunctionParameterDecl(
name='dstsig',
type=ValueType(name='unsigned int'),
),
),
doc='Extract a subset of components from a state previously obtained via mj_getState.', # pylint: disable=line-too-long
)),
('mj_setState',
FunctionDecl(
name='mj_setState',
+22
View File
@@ -211,6 +211,28 @@ void mj_getState(const mjModel* m, const mjData* d, mjtNum* state, unsigned int
}
// extract a sub-state from a state
void mj_extractState(const mjModel* m, const mjtNum* src, unsigned int srcsig,
mjtNum* dst, unsigned int dstsig) {
if ((srcsig & dstsig) != dstsig) {
mjERROR("dstsig is not a subset of srcsig");
return;
}
for (int i=0; i < mjNSTATE; i++) {
mjtState element = 1<<i;
if (element & srcsig) {
int size = mj_stateElemSize(m, element);
if (element & dstsig) {
mju_copy(dst, src, size);
dst += size;
}
src += size;
}
}
}
// set state
void mj_setState(const mjModel* m, mjData* d, const mjtNum* state, unsigned int sig) {
if (sig >= (1<<mjNSTATE)) {
+4
View File
@@ -41,6 +41,10 @@ MJAPI int mj_stateSize(const mjModel* m, unsigned int sig);
// get state
MJAPI void mj_getState(const mjModel* m, const mjData* d, mjtNum* state, unsigned int sig);
// extract a sub-state from a state
MJAPI void mj_extractState(const mjModel* m, const mjtNum* src, unsigned int srcsig,
mjtNum* dst, unsigned int dstsig);
// set state
MJAPI void mj_setState(const mjModel* m, mjData* d, const mjtNum* state, unsigned int sig);
+68
View File
@@ -17,9 +17,11 @@
#include "src/engine/engine_core_util.h"
#include "src/engine/engine_support.h"
#include <cstring>
#include <limits>
#include <random>
#include <string>
#include <string_view>
#include <vector>
#include <gmock/gmock.h>
@@ -742,6 +744,72 @@ TEST_F(SupportTest, GetSetStateStepEqual) {
mj_deleteModel(model);
}
TEST_F(SupportTest, ExtractState) {
const std::string xml_path = GetTestDataFilePath(kDefaultModel);
mjModel* model = mj_loadXML(xml_path.c_str(), nullptr, nullptr, 0);
mjData* data = mj_makeData(model);
// make distribution using seed
std::mt19937_64 rng;
rng.seed(3);
std::normal_distribution<double> dist(0, .01);
// set controls and applied joint forces to random values
for (int i=0; i < model->nu; i++) data->ctrl[i] = dist(rng);
for (int i=0; i < model->nv; i++) data->qfrc_applied[i] = dist(rng);
for (int i=0; i < model->neq; i++) data->eq_active[i] = dist(rng) > 0;
// take one step
mj_step(model, data);
// take a state that will be used as src
int srcsig = mjSTATE_TIME | mjSTATE_QPOS | mjSTATE_QVEL | mjSTATE_CTRL;
int srcsize = mj_stateSize(model, srcsig);
vector<mjtNum> srcstate(srcsize);
mj_getState(model, data, srcstate.data(), srcsig);
// extract a subset consisting of only a single bit in srcsig
int dstsig1 = mjSTATE_CTRL;
int dstsize1 = mj_stateSize(model, dstsig1);
EXPECT_LT(dstsize1, srcsize);
EXPECT_EQ(dstsize1, model->nu);
vector<mjtNum> dststate1(dstsize1);
mj_extractState(model, srcstate.data(), srcsig, dststate1.data(), dstsig1);
EXPECT_EQ(dststate1, AsVector(data->ctrl, model->nu));
// extract a subset consisting of multiple non-consecutive bits in srcsig
int dstsig2 = mjSTATE_QPOS | mjSTATE_CTRL;
int dstsize2 = mj_stateSize(model, dstsig2);
EXPECT_LT(dstsize2, srcsize);
EXPECT_EQ(dstsize2, model->nq + model->nu);
vector<mjtNum> dststate2(dstsize2);
mj_extractState(model, srcstate.data(), srcsig, dststate2.data(), dstsig2);
EXPECT_EQ(AsVector(dststate2.data(), model->nq),
AsVector(data->qpos, model->nq));
EXPECT_EQ(AsVector(dststate2.data() + model->nq, model->nu),
AsVector(data->ctrl, model->nu));
// test that an error is correctly raised if dstsig is not a subset of srcsig
static int error_count;
static char last_error_msg[128];
error_count = 0;
last_error_msg[0] = '\0';
auto* error_handler = +[](const char* msg) {
std::strncpy(last_error_msg, msg, sizeof(last_error_msg));
++error_count;
};
auto* old_mju_user_error = mju_user_error;
mju_user_error = error_handler;
mj_extractState(model, nullptr, srcsig, nullptr, mjSTATE_QFRC_APPLIED);
mju_user_error = old_mju_user_error;
EXPECT_EQ(error_count, 1);
EXPECT_EQ(std::string_view(last_error_msg),
"mj_extractState: dstsig is not a subset of srcsig");
mj_deleteData(data);
mj_deleteModel(model);
}
using InertiaTest = MujocoTest;
static const char* const kInertiaPath = "engine/testdata/inertia.xml";
+3
View File
@@ -6621,6 +6621,9 @@ public static unsafe extern int mj_stateSize(mjModel_* m, uint sig);
[DllImport("mujoco", CallingConvention = CallingConvention.Cdecl)]
public static unsafe extern void mj_getState(mjModel_* m, mjData_* d, double* state, uint sig);
[DllImport("mujoco", CallingConvention = CallingConvention.Cdecl)]
public static unsafe extern void mj_extractState(mjModel_* m, double* src, uint srcsig, double* dst, uint dstsig);
[DllImport("mujoco", CallingConvention = CallingConvention.Cdecl)]
public static unsafe extern void mj_setState(mjModel_* m, mjData_* d, double* state, uint sig);