Resizes keyframes before storing them during attach/detach.

Improve keyframe tests.

PiperOrigin-RevId: 685450257
Change-Id: I97aa129c4d7a66f613290605b85657e8081700b6
This commit is contained in:
Alessio Quaglino
2024-10-13 10:01:26 -07:00
committed by Copybara-Service
parent 534aedd47d
commit 52a0149cd1
6 changed files with 199 additions and 67 deletions
+6 -1
View File
@@ -198,7 +198,12 @@ const char* mjs_getError(mjSpec* s) {
int mjs_detachBody(mjSpec* s, mjsBody* b) {
mjCModel* model = static_cast<mjCModel*>(s->element);
mjCBody* body = static_cast<mjCBody*>(b->element);
*model -= *body;
try {
*model -= *body;
} catch (mjCError& e) {
model->SetError(e);
return -1;
}
delete body;
return 0;
}
+112 -52
View File
@@ -264,11 +264,11 @@ void mjCModel::ResetTreeLists() {
// save associated state addresses in related elements
void mjCModel::SaveDofOffsets() {
void mjCModel::SaveDofOffsets(bool computesize) {
int qposadr = 0;
int dofadr = 0;
int actadr = 0;
int nmocap = 0;
int mocapadr = 0;
for (auto joint : joints_) {
joint->qposadr_ = qposadr;
@@ -289,12 +289,19 @@ void mjCModel::SaveDofOffsets() {
for (mjCBody* body : bodies_) {
if (body->spec.mocap) {
body->mocapid = nmocap;
nmocap++;
body->mocapid = mocapadr++;
} else {
body->mocapid = -1;
}
}
if (computesize) {
nq = qposadr;
nv = dofadr;
na = actadr;
nu = (int)actuators_.size();
nmocap = mocapadr;
}
}
template <class T>
@@ -1556,6 +1563,7 @@ void mjCModel::SetSizes() {
ntuple = (int)tuples_.size();
nkey = (int)keys_.size();
nplugin = (int)plugins_.size();
nq = nv = nu = na = nmocap = 0;
// nq, nv
for (int i=0; i<njnt; i++) {
@@ -3015,7 +3023,11 @@ void mjCModel::CopyObjects(mjModel* m) {
// save qpos0 in user model (to recognize changed key_qpos in write)
qpos0.resize(nq);
body_pos0.resize(3*nbody);
body_quat0.resize(4*nbody);
mjuu_copyvec(qpos0.data(), m->qpos0, nq);
mjuu_copyvec(body_pos0.data(), m->body_pos, 3*nbody);
mjuu_copyvec(body_quat0.data(), m->body_quat, 4*nbody);
}
@@ -3101,7 +3113,7 @@ void mjCModel::RestoreState(const std::string& state_name, const mjtNum* pos0,
for (unsigned int i=0; i<bodies_.size(); i++) {
auto body = bodies_[i];
if (!body->mocap) {
if (!body->spec.mocap) {
continue;
}
if (mpos) {
@@ -3143,11 +3155,16 @@ void mjCModel::StoreKeyframes() {
resetlists = true;
}
// do not change the offset computed during compilation in case the user wants to recompile
// do not change compilation quantities in case the user wants to recompile preserving the state
if (!compiled) {
SaveDofOffsets();
SaveDofOffsets(/*computesize=*/true);
qpos0.resize(nq);
body_pos0.resize(3*bodies_.size());
body_quat0.resize(4*bodies_.size());
ComputeReference(qpos0, body_pos0, body_quat0);
}
// save keyframe info and resize keyframes
for (auto& key : keys_) {
mjKeyInfo info;
info.name = prefix + key->name + suffix;
@@ -3159,6 +3176,7 @@ void mjCModel::StoreKeyframes() {
info.mpos = !key->spec_mpos_.empty();
info.mquat = !key->spec_mquat_.empty();
key_pending_.push_back(info);
ResizeKeyframe(key, qpos0.data(), body_pos0.data(), body_quat0.data());
SaveState(info.name, key->spec_qpos_.data(), key->spec_qvel_.data(),
key->spec_act_.data(), key->spec_ctrl_.data(),
key->spec_mpos_.data(), key->spec_mquat_.data());
@@ -3167,6 +3185,10 @@ void mjCModel::StoreKeyframes() {
if (resetlists) {
ResetTreeLists();
}
if (!compiled) {
nq = nv = na = nu = nmocap = 0;
}
}
@@ -3697,6 +3719,87 @@ void mjCModel::CompileMeshes(const mjVFS* vfs) {
// compute qpos0
template <class T>
void mjCModel::ComputeReference(std::vector<T>& q0, std::vector<T>& bpos, std::vector<T>& bquat) {
int b = 0;
for (auto body : bodies_) {
mjuu_copyvec(bpos.data()+3*b, body->spec.pos, 3);
mjuu_copyvec(bquat.data()+4*b, body->spec.quat, 4);
for (auto joint : body->joints) {
switch (joint->type) {
case mjJNT_FREE:
mjuu_copyvec(q0.data()+joint->qposadr_, body->spec.pos, 3);
mjuu_copyvec(q0.data()+joint->qposadr_+3, body->spec.quat, 4);
break;
case mjJNT_BALL:
mjuu_setvec(q0.data()+joint->qposadr_, 1, 0, 0, 0);
break;
case mjJNT_SLIDE:
case mjJNT_HINGE:
q0[joint->qposadr_] = (T)joint->spec.ref;
break;
default:
throw mjCError(joint, "unknown joint type");
}
}
b++;
}
}
// resizes a keyframe, filling in missing values
void mjCModel::ResizeKeyframe(mjCKey* key, const mjtNum* qpos0_,
const mjtNum* bpos, const mjtNum* bquat) {
if (!key->spec_qpos_.empty()) {
int nq0 = key->spec_qpos_.size();
key->spec_qpos_.resize(nq);
for (int i=nq0; i<nq; i++) {
key->spec_qpos_[i] = (double)qpos0_[i];
}
}
if (!key->spec_qvel_.empty()) {
key->spec_qvel_.resize(nv);
}
if (!key->spec_act_.empty()) {
key->spec_act_.resize(na);
}
if (!key->spec_ctrl_.empty()) {
key->spec_ctrl_.resize(nu);
}
if (!key->spec_mpos_.empty()) {
int nmocap0 = key->spec_mpos_.size() / 3;
key->spec_mpos_.resize(3*nmocap);
for (unsigned int j = 0; j < bodies_.size(); j++) {
if (bodies_[j]->mocapid < nmocap0) {
continue;
}
int i = bodies_[j]->mocapid;
key->spec_mpos_[3*i+0] = (double)bpos[3*j+0];
key->spec_mpos_[3*i+1] = (double)bpos[3*j+1];
key->spec_mpos_[3*i+2] = (double)bpos[3*j+2];
}
}
if (!key->spec_mquat_.empty()) {
int nmocap0 = key->spec_mquat_.size() / 4;
key->spec_mquat_.resize(4*nmocap);
for (unsigned int j = 0; j < bodies_.size(); j++) {
if (bodies_[j]->mocapid < nmocap0) {
continue;
}
int i = bodies_[j]->mocapid;
key->spec_mquat_[4*i+0] = (double)bquat[4*j+0];
key->spec_mquat_[4*i+1] = (double)bquat[4*j+1];
key->spec_mquat_[4*i+2] = (double)bquat[4*j+2];
key->spec_mquat_[4*i+3] = (double)bquat[4*j+3];
}
}
}
// convert pending keyframes info to actual keyframes
void mjCModel::ResolveKeyframes(const mjModel* m) {
if (key_pending_.empty()) {
@@ -3707,51 +3810,8 @@ void mjCModel::ResolveKeyframes(const mjModel* m) {
SaveDofOffsets();
// resize existing keyframes to the new state, fill in missing default values
for (unsigned int i = 0; i < nkey - key_pending_.size(); i++) {
mjCKey* key = keys_[i];
if (!key->spec_qpos_.empty()) {
int nq0 = key->spec_qpos_.size();
key->spec_qpos_.resize(nq);
for (int i=nq0; i<m->nq; i++) {
key->spec_qpos_[i] = (double)m->qpos0[i];
}
}
if (!key->spec_qvel_.empty()) {
key->spec_qvel_.resize(nv);
}
if (!key->spec_act_.empty()) {
key->spec_act_.resize(na);
}
if (!key->spec_ctrl_.empty()) {
key->spec_ctrl_.resize(nu);
}
if (!key->spec_mpos_.empty()) {
int nmocap0 = key->spec_mpos_.size() / 3;
key->spec_mpos_.resize(3*nmocap);
for (unsigned int j = 0; j < bodies_.size(); j++) {
if (bodies_[j]->mocapid < nmocap0) {
continue;
}
int i = bodies_[j]->mocapid;
key->spec_mpos_[3*i+0] = (double)m->body_pos[3*j+0];
key->spec_mpos_[3*i+1] = (double)m->body_pos[3*j+1];
key->spec_mpos_[3*i+2] = (double)m->body_pos[3*j+2];
}
}
if (!key->spec_mquat_.empty()) {
int nmocap0 = key->spec_mquat_.size() / 4;
key->spec_mquat_.resize(4*nmocap);
for (unsigned int j = 0; j < bodies_.size(); j++) {
if (bodies_[j]->mocapid < nmocap0) {
continue;
}
int i = bodies_[j]->mocapid;
key->spec_mquat_[4*i+0] = (double)m->body_quat[4*j+0];
key->spec_mquat_[4*i+1] = (double)m->body_quat[4*j+1];
key->spec_mquat_[4*i+2] = (double)m->body_quat[4*j+2];
key->spec_mquat_[4*i+3] = (double)m->body_quat[4*j+3];
}
}
for (auto* key : keys_) {
ResizeKeyframe(key, m->qpos0, m->body_pos, m->body_quat);
}
// create new keyframes, fill in missing default values
+11 -1
View File
@@ -128,6 +128,8 @@ class mjCModel_ : public mjsElement {
// save qpos0, to recognize changed key_qpos in write
std::vector<mjtNum> qpos0;
std::vector<mjtNum> body_pos0;
std::vector<mjtNum> body_quat0;
// variable-size attributes
std::string comment_; // comment at top of XML
@@ -382,11 +384,19 @@ class mjCModel : public mjCModel_, private mjSpec {
void ResetTreeLists();
// save dof offsets in joints and actuators
void SaveDofOffsets();
void SaveDofOffsets(bool computesize = false);
// convert pending keyframes info to actual keyframes
void ResolveKeyframes(const mjModel* m);
// resize a keyframe, filling in missing values
void ResizeKeyframe(mjCKey* key, const mjtNum* qpos0_, const mjtNum* bpos, const mjtNum* bquat);
// compute qpos0
template <class T>
void ComputeReference(std::vector<T>& q0, std::vector<T>& bpos,
std::vector<T>& bquat);
mjListKeyMap ids; // map from object names to ids
mjCError errInfo; // last error info
std::vector<mjKeyInfo> key_pending_; // attached keyframes
+3 -1
View File
@@ -3617,7 +3617,9 @@ void mjXReader::Body(XMLElement* section, mjsBody* body, mjsFrame* frame,
}
// delete subtree
mjs_detachBody(spec, subtree);
if (mjs_detachBody(spec, subtree)) {
throw mjXError(elem, mjs_getError(spec));
}
}
// body sub-element
+43
View File
@@ -783,6 +783,10 @@ TEST_F(MujocoTest, AttachDifferent) {
<frame name="frame" pos=".1 0 0" euler="0 90 0"/>
</body>
</worldbody>
<keyframe>
<key name="one" time="1" qpos="1 1 1 1 0 0 0"/>
</keyframe>
</mujoco>)";
static constexpr char xml_result[] = R"(
@@ -837,6 +841,7 @@ TEST_F(MujocoTest, AttachDifferent) {
</contact>
<keyframe>
<key name="one" time="1" qpos="1 1 1 1 0 0 0 0"/>
<key name="attached-two-1" time="2" qpos="0 0 0 1 0 0 0 2" act="2 2" ctrl="2 2"/>
<key name="attached-three-1" time="3" qpos="0 0 0 1 0 0 0 3" act="3 3" ctrl="3 3"/>
</keyframe>
@@ -905,6 +910,10 @@ TEST_F(MujocoTest, AttachFrame) {
<frame name="frame" pos=".1 0 0" euler="0 90 0"/>
</body>
</worldbody>
<keyframe>
<key name="one" time="1" qpos="1 1 1 1 0 0 0"/>
</keyframe>
</mujoco>)";
static constexpr char xml_result[] = R"(
@@ -959,6 +968,7 @@ TEST_F(MujocoTest, AttachFrame) {
</contact>
<keyframe>
<key name="one" time="1" qpos="1 1 1 1 0 0 0 0"/>
<key name="attached-two-1" time="2" qpos="0 0 0 1 0 0 0 2" act="2 2" ctrl="2 2"/>
<key name="attached-three-1" time="3" qpos="0 0 0 1 0 0 0 3" act="3 3" ctrl="3 3"/>
</keyframe>
@@ -1356,6 +1366,39 @@ TEST_F(MujocoTest, AttachMocap) {
mj_deleteModel(m_expected);
}
TEST_F(MujocoTest, ReplicateKeyframe) {
static constexpr char xml[] = R"(
<mujoco>
<worldbody>
<replicate count="1" euler="0 0 1.8">
<body name="body" pos="0 -1 0">
<joint type="slide"/>
<geom name="g" size="1"/>
</body>
</replicate>
</worldbody>
<keyframe>
<key name="keyframe" qpos="1"/>
</keyframe>
</mujoco>
)";
std::array<char, 1024> error;
mjModel* m = LoadModelFromString(xml, error.data(), error.size());
EXPECT_THAT(m, testing::NotNull()) << error.data();
EXPECT_THAT(m->ngeom, 1);
EXPECT_THAT(m->nbody, 2);
// check that the keyframe is resized
EXPECT_THAT(m->nkey, 1);
EXPECT_THAT(m->nq, 1);
EXPECT_THAT(m->key_qpos[0], 0);
EXPECT_STREQ(mj_id2name(m, mjOBJ_KEY, 0), "keyframe");
mj_deleteModel(m);
}
TEST_F(MujocoTest, AttachUnnamedAssets) {
static constexpr char cube[] = R"(
v -1 -1 1
+24 -12
View File
@@ -1234,21 +1234,28 @@ TEST_F(XMLReaderTest, ParseReplicate) {
</asset>
<worldbody>
<replicate count="101" euler="0 0 1.8">
<body name="body" pos="0 -1 0">
<joint type="slide"/>
<geom name="g" size="1"/>
</body>
</replicate>
<replicate count="2" offset="1 0 0">
<replicate count="2" offset="0 1 0" sep="_">
<geom name="geom" size="1" pos="0 0 1" material="material"/>
<site name="site" pos="1 0 0"/>
</replicate>
</replicate>
<replicate count="101" euler="0 0 1.8">
<geom name="g" size="1" pos="0 -1 0"/>
</replicate>
</worldbody>
<sensor>
<framepos name="sensor" objtype="site" objname="site"/>
</sensor>
<keyframe>
<key name="keyframe" qpos="1"/>
</keyframe>
</mujoco>
)";
@@ -1287,14 +1294,19 @@ TEST_F(XMLReaderTest, ParseReplicate) {
}
// check that the final pose is correct
int n = 104;
EXPECT_NEAR(m->geom_pos[3*n+0], 0, 1e-8);
EXPECT_NEAR(m->geom_pos[3*n+1], 1, 1e-8);
EXPECT_EQ(m->geom_pos[3*n+2], 0);
EXPECT_NEAR(m->geom_quat[4*n+0], 0, 1e-8);
EXPECT_EQ(m->geom_quat[4*n+1], 0);
EXPECT_EQ(m->geom_quat[4*n+2], 0);
EXPECT_EQ(m->geom_quat[4*n+3], 1);
int n = m->nbody-1;
EXPECT_THAT(m->nbody, 102);
EXPECT_NEAR(m->body_pos[3*n+0], 0, 1e-8);
EXPECT_NEAR(m->body_pos[3*n+1], 1, 1e-8);
EXPECT_EQ(m->body_pos[3*n+2], 0);
EXPECT_NEAR(m->body_quat[4*n+0], 0, 1e-8);
EXPECT_EQ(m->body_quat[4*n+1], 0);
EXPECT_EQ(m->body_quat[4*n+2], 0);
EXPECT_EQ(m->body_quat[4*n+3], 1);
// check that the pending keyframes are lost while detaching
EXPECT_THAT(m->nkey, 0);
EXPECT_THAT(m->nq, 101);
mj_deleteModel(m);
}