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
+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";