From 57241585934cd646c8f99eca400577dc57c791aa Mon Sep 17 00:00:00 2001 From: Alessio Quaglino Date: Tue, 14 Apr 2026 09:38:17 -0700 Subject: [PATCH] Refactor flex stiffness storage to support variable sizes. This change introduces `nflexstiffness` and `flex_stiffnessadr` to allow flex stiffness matrices to have sizes other than the fixed 21 per element. Higher-order flexes can now store larger stiffness matrices based on their number of nodes. The `flex_stiffnessadr` array provides the starting index for each flex's stiffness data within the `flex_stiffness` array. PiperOrigin-RevId: 899632263 Change-Id: Ie49182c46c3777acf0c6b492345edfcf62fb5e44 --- doc/includes/references.h | 4 ++- include/mujoco/mjmodel.h | 4 ++- include/mujoco/mjxmacro.h | 4 ++- python/mujoco/introspect/structs.py | 15 +++++++++- src/engine/engine_derivative.c | 2 +- src/engine/engine_io.c | 42 +++++++++++++--------------- src/engine/engine_io.h | 2 +- src/engine/engine_passive.c | 2 +- src/user/user_mesh.cc | 13 +++++---- src/user/user_model.cc | 36 ++++++++++++++++++------ src/user/user_model.h | 1 + test/user/user_flex_test.cc | 8 ++++-- unity/Runtime/Bindings/MjBindings.cs | 2 ++ wasm/codegen/generated/bindings.cc | 13 ++++++++- 14 files changed, 100 insertions(+), 48 deletions(-) diff --git a/doc/includes/references.h b/doc/includes/references.h index bea7f4d9..7f23a52d 100644 --- a/doc/includes/references.h +++ b/doc/includes/references.h @@ -1040,6 +1040,7 @@ struct mjModel_ { mjtSize nflexedge; // number of edges in all flexes mjtSize nflexelem; // number of elements in all flexes mjtSize nflexelemdata; // number of element vertex ids in all flexes + mjtSize nflexstiffness; // number of stiffness parameters in all flexes mjtSize nflexelemedge; // number of element edge ids in all flexes mjtSize nflexshelldata; // number of shell fragment vertex ids in all flexes mjtSize nflexevpair; // number of element-vertex pairs in all flexes @@ -1323,6 +1324,7 @@ struct mjModel_ { int* flex_elemadr; // first element address (nflex x 1) int* flex_elemnum; // number of elements (nflex x 1) int* flex_elemdataadr; // first element vertex id address (nflex x 1) + int* flex_stiffnessadr; // stiffness matrix address (nflex x 1) int* flex_elemedgeadr; // first element edge id address (nflex x 1) int* flex_shellnum; // number of shells (nflex x 1) int* flex_shelldataadr; // first shell data address (nflex x 1) @@ -1351,7 +1353,7 @@ struct mjModel_ { mjtNum* flexedge_invweight0; // edge inv. weight in qpos0 (nflexedge x 1) mjtNum* flex_radius; // radius around primitive element (nflex x 1) mjtNum* flex_size; // vertex bounding box half sizes in qpos0 (nflex x 3) - mjtNum* flex_stiffness; // finite element stiffness matrix (nflexelem x 21) + mjtNum* flex_stiffness; // finite element stiffness matrix (nflexstiffness x 1) mjtNum* flex_bending; // bending stiffness (nflexedge x 17) mjtNum* flex_damping; // Rayleigh's damping coefficient (nflex x 1) mjtNum* flex_edgestiffness; // edge stiffness (nflex x 1) diff --git a/include/mujoco/mjmodel.h b/include/mujoco/mjmodel.h index 49cfee0b..8a45f67b 100644 --- a/include/mujoco/mjmodel.h +++ b/include/mujoco/mjmodel.h @@ -703,6 +703,7 @@ struct mjModel_ { mjtSize nflexedge; // number of edges in all flexes mjtSize nflexelem; // number of elements in all flexes mjtSize nflexelemdata; // number of element vertex ids in all flexes + mjtSize nflexstiffness; // number of stiffness parameters in all flexes mjtSize nflexelemedge; // number of element edge ids in all flexes mjtSize nflexshelldata; // number of shell fragment vertex ids in all flexes mjtSize nflexevpair; // number of element-vertex pairs in all flexes @@ -986,6 +987,7 @@ struct mjModel_ { int* flex_elemadr; // first element address (nflex x 1) int* flex_elemnum; // number of elements (nflex x 1) int* flex_elemdataadr; // first element vertex id address (nflex x 1) + int* flex_stiffnessadr; // stiffness matrix address (nflex x 1) int* flex_elemedgeadr; // first element edge id address (nflex x 1) int* flex_shellnum; // number of shells (nflex x 1) int* flex_shelldataadr; // first shell data address (nflex x 1) @@ -1014,7 +1016,7 @@ struct mjModel_ { mjtNum* flexedge_invweight0; // edge inv. weight in qpos0 (nflexedge x 1) mjtNum* flex_radius; // radius around primitive element (nflex x 1) mjtNum* flex_size; // vertex bounding box half sizes in qpos0 (nflex x 3) - mjtNum* flex_stiffness; // finite element stiffness matrix (nflexelem x 21) + mjtNum* flex_stiffness; // finite element stiffness matrix (nflexstiffness x 1) mjtNum* flex_bending; // bending stiffness (nflexedge x 17) mjtNum* flex_damping; // Rayleigh's damping coefficient (nflex x 1) mjtNum* flex_edgestiffness; // edge stiffness (nflex x 1) diff --git a/include/mujoco/mjxmacro.h b/include/mujoco/mjxmacro.h index 318e76da..970ba622 100644 --- a/include/mujoco/mjxmacro.h +++ b/include/mujoco/mjxmacro.h @@ -185,6 +185,7 @@ X( nflexedge ) \ X( nflexelem ) \ X( nflexelemdata ) \ + X( nflexstiffness ) \ X( nflexelemedge ) \ X( nflexshelldata ) \ X( nflexevpair ) \ @@ -462,6 +463,7 @@ X ( int, flex_elemadr, nflex, 1 ) \ X ( int, flex_elemnum, nflex, 1 ) \ X ( int, flex_elemdataadr, nflex, 1 ) \ + X ( int, flex_stiffnessadr, nflex, 1 ) \ X ( int, flex_elemedgeadr, nflex, 1 ) \ X ( int, flex_shellnum, nflex, 1 ) \ X ( int, flex_shelldataadr, nflex, 1 ) \ @@ -490,7 +492,7 @@ X ( mjtNum, flexedge_invweight0, nflexedge, 1 ) \ X ( mjtNum, flex_radius, nflex, 1 ) \ X ( mjtNum, flex_size, nflex, 3 ) \ - X ( mjtNum, flex_stiffness, nflexelem, 21 ) \ + X ( mjtNum, flex_stiffness, nflexstiffness, 1 ) \ X ( mjtNum, flex_bending, nflexedge, 17 ) \ X ( mjtNum, flex_damping, nflex, 1 ) \ X ( mjtNum, flex_edgestiffness, nflex, 1 ) \ diff --git a/python/mujoco/introspect/structs.py b/python/mujoco/introspect/structs.py index 57a0476e..4dded269 100644 --- a/python/mujoco/introspect/structs.py +++ b/python/mujoco/introspect/structs.py @@ -972,6 +972,11 @@ STRUCTS: Mapping[str, StructDecl] = dict([ type=ValueType(name='mjtSize'), doc='number of element vertex ids in all flexes', ), + StructFieldDecl( + name='nflexstiffness', + type=ValueType(name='mjtSize'), + doc='number of stiffness parameters in all flexes', + ), StructFieldDecl( name='nflexelemedge', type=ValueType(name='mjtSize'), @@ -2735,6 +2740,14 @@ STRUCTS: Mapping[str, StructDecl] = dict([ doc='first element vertex id address', array_extent=('nflex',), ), + StructFieldDecl( + name='flex_stiffnessadr', + type=PointerType( + inner_type=ValueType(name='int'), + ), + doc='stiffness matrix address', + array_extent=('nflex',), + ), StructFieldDecl( name='flex_elemedgeadr', type=PointerType( @@ -2965,7 +2978,7 @@ STRUCTS: Mapping[str, StructDecl] = dict([ inner_type=ValueType(name='mjtNum'), ), doc='finite element stiffness matrix', - array_extent=('nflexelem', 21), + array_extent=('nflexstiffness',), ), StructFieldDecl( name='flex_bending', diff --git a/src/engine/engine_derivative.c b/src/engine/engine_derivative.c index 23781e1d..5f1058d0 100644 --- a/src/engine/engine_derivative.c +++ b/src/engine/engine_derivative.c @@ -899,7 +899,7 @@ static void mjd_flexInterp_kernel(const mjModel* m, mjData* d, mjtFlexOp op, } // get stiffness and damping - mjtNum* k = m->flex_stiffness + 21*m->flex_elemadr[f]; + mjtNum* k = m->flex_stiffness + m->flex_stiffnessadr[f]; // skip if rigid or no stiffness if (m->flex_rigid[f] || k[0] == 0) { diff --git a/src/engine/engine_io.c b/src/engine/engine_io.c index 8a8e59b0..33fafe6c 100644 --- a/src/engine/engine_io.c +++ b/src/engine/engine_io.c @@ -205,11 +205,11 @@ void mj_makeModel(mjModel** dest, mjtSize nbvhdynamic, mjtSize noct, mjtSize njnt, mjtSize ntree, mjtSize nM, mjtSize nB, mjtSize nC, mjtSize nD, mjtSize ngeom, mjtSize nsite, mjtSize ncam, mjtSize nlight, mjtSize nflex, mjtSize nflexnode, mjtSize nflexvert, mjtSize nflexedge, mjtSize nflexelem, - mjtSize nflexelemdata, mjtSize nflexelemedge, mjtSize nflexshelldata, mjtSize nflexevpair, - mjtSize nflextexcoord, mjtSize nJfe, mjtSize nJfv, mjtSize nmesh, mjtSize nmeshvert, - mjtSize nmeshnormal, mjtSize nmeshtexcoord, mjtSize nmeshface, mjtSize nmeshgraph, - mjtSize nmeshpoly, mjtSize nmeshpolyvert, mjtSize nmeshpolymap, mjtSize nskin, - mjtSize nskinvert, mjtSize nskintexvert, mjtSize nskinface, mjtSize nskinbone, + mjtSize nflexelemdata, mjtSize nflexstiffness, mjtSize nflexelemedge, mjtSize nflexshelldata, + mjtSize nflexevpair, mjtSize nflextexcoord, mjtSize nJfe, mjtSize nJfv, mjtSize nmesh, + mjtSize nmeshvert, mjtSize nmeshnormal, mjtSize nmeshtexcoord, mjtSize nmeshface, + mjtSize nmeshgraph, mjtSize nmeshpoly, mjtSize nmeshpolyvert, mjtSize nmeshpolymap, + mjtSize nskin, mjtSize nskinvert, mjtSize nskintexvert, mjtSize nskinface, mjtSize nskinbone, mjtSize nskinbonevert, mjtSize nhfield, mjtSize nhfielddata, mjtSize ntex, mjtSize ntexdata, mjtSize nmat, mjtSize npair, mjtSize nexclude, mjtSize neq, mjtSize ntendon, mjtSize nJten, mjtSize nwrap, mjtSize nsensor, mjtSize nnumeric, mjtSize nnumericdata, mjtSize ntext, @@ -293,6 +293,7 @@ void mj_makeModel(mjModel** dest, m->nflexedge = nflexedge; m->nflexelem = nflexelem; m->nflexelemdata = nflexelemdata; + m->nflexstiffness = nflexstiffness; m->nflexelemedge = nflexelemedge; m->nflexshelldata = nflexshelldata; m->nflexevpair = nflexevpair; @@ -400,22 +401,19 @@ mjModel* mj_copyModel(mjModel* dest, const mjModel* src) { // allocate new model if needed if (!dest) { mj_makeModel( - &dest, src->nq, src->nv, src->nu, src->na, src->nbody, src->nbvh, - src->nbvhstatic, src->nbvhdynamic, src->noct, src->njnt, src->ntree, - src->nM, src->nB, src->nC, src->nD, src->ngeom, src->nsite, src->ncam, - src->nlight, src->nflex, src->nflexnode, src->nflexvert, src->nflexedge, - src->nflexelem, src->nflexelemdata, src->nflexelemedge, src->nflexshelldata, - src->nflexevpair, src->nflextexcoord, src->nJfe, src->nJfv, src->nmesh, - src->nmeshvert, src->nmeshnormal, src->nmeshtexcoord, src->nmeshface, - src->nmeshgraph, src->nmeshpoly, src->nmeshpolyvert, src->nmeshpolymap, - src->nskin, src->nskinvert, src->nskintexvert, src->nskinface, - src->nskinbone, src->nskinbonevert, src->nhfield, src->nhfielddata, - src->ntex, src->ntexdata, src->nmat, src->npair, src->nexclude, - src->neq, src->ntendon, src->nJten, src->nwrap, src->nsensor, - src->nnumeric, src->nnumericdata, src->ntext, src->ntextdata, - src->ntuple, src->ntupledata, src->nkey, src->nmocap, src->nplugin, - src->npluginattr, src->nuser_body, src->nuser_jnt, src->nuser_geom, - src->nuser_site, src->nuser_cam, src->nuser_tendon, src->nuser_actuator, + &dest, src->nq, src->nv, src->nu, src->na, src->nbody, src->nbvh, src->nbvhstatic, + src->nbvhdynamic, src->noct, src->njnt, src->ntree, src->nM, src->nB, src->nC, src->nD, + src->ngeom, src->nsite, src->ncam, src->nlight, src->nflex, src->nflexnode, src->nflexvert, + src->nflexedge, src->nflexelem, src->nflexelemdata, src->nflexstiffness, + src->nflexelemedge, src->nflexshelldata, src->nflexevpair, src->nflextexcoord, src->nJfe, + src->nJfv, src->nmesh, src->nmeshvert, src->nmeshnormal, src->nmeshtexcoord, src->nmeshface, + src->nmeshgraph, src->nmeshpoly, src->nmeshpolyvert, src->nmeshpolymap, src->nskin, + src->nskinvert, src->nskintexvert, src->nskinface, src->nskinbone, src->nskinbonevert, + src->nhfield, src->nhfielddata, src->ntex, src->ntexdata, src->nmat, src->npair, + src->nexclude, src->neq, src->ntendon, src->nJten, src->nwrap, src->nsensor, src->nnumeric, + src->nnumericdata, src->ntext, src->ntextdata, src->ntuple, src->ntupledata, src->nkey, + src->nmocap, src->nplugin, src->npluginattr, src->nuser_body, src->nuser_jnt, + src->nuser_geom, src->nuser_site, src->nuser_cam, src->nuser_tendon, src->nuser_actuator, src->nuser_sensor, src->nnames, src->npaths); } if (!dest) { @@ -599,7 +597,7 @@ mjModel* mj_loadModelBuffer(const void* buffer, int buffer_sz) { sizes[56], sizes[57], sizes[58], sizes[59], sizes[60], sizes[61], sizes[62], sizes[63], sizes[64], sizes[65], sizes[66], sizes[67], sizes[68], sizes[69], sizes[70], sizes[71], sizes[72], sizes[73], sizes[74], sizes[75], sizes[76], - sizes[77]); + sizes[77], sizes[78]); // mj_makeModel may fail if the input buffer has invalid sizes if (!m) { diff --git a/src/engine/engine_io.h b/src/engine/engine_io.h index 0af1058f..1aedaf86 100644 --- a/src/engine/engine_io.h +++ b/src/engine/engine_io.h @@ -52,7 +52,7 @@ void mj_makeModel(mjModel** dest, mjtSize nbvhdynamic, mjtSize noct, mjtSize njnt, mjtSize ntree, mjtSize nM, mjtSize nB, mjtSize nC, mjtSize nD, mjtSize ngeom, mjtSize nsite, mjtSize ncam, mjtSize nlight, mjtSize nflex, mjtSize nflexnode, mjtSize nflexvert, mjtSize nflexedge, mjtSize nflexelem, - mjtSize nflexelemdata, mjtSize nflexelemedge, mjtSize nflexshelldata, mjtSize nflexevpair, + mjtSize nflexelemdata, mjtSize nflexstiffness, mjtSize nflexelemedge, mjtSize nflexshelldata, mjtSize nflexevpair, mjtSize nflextexcoord, mjtSize nJfe, mjtSize nJfv, mjtSize nmesh, mjtSize nmeshvert, mjtSize nmeshnormal, mjtSize nmeshtexcoord, mjtSize nmeshface, mjtSize nmeshgraph, mjtSize nmeshpoly, mjtSize nmeshpolyvert, mjtSize nmeshpolymap, mjtSize nskin, diff --git a/src/engine/engine_passive.c b/src/engine/engine_passive.c index d27f615b..d67f8683 100644 --- a/src/engine/engine_passive.c +++ b/src/engine/engine_passive.c @@ -201,7 +201,7 @@ static void mj_springdamper(const mjModel* m, mjData* d) { // flex elasticity for (int f=0; f < m->nflex; f++) { - mjtNum* k = m->flex_stiffness + 21*m->flex_elemadr[f]; + mjtNum* k = m->flex_stiffness + m->flex_stiffnessadr[f]; mjtNum* b = m->flex_bending + 17*m->flex_edgeadr[f]; int dim = m->flex_dim[f]; int nodenum = m->flex_nodenum[f]; diff --git a/src/user/user_mesh.cc b/src/user/user_mesh.cc index db852729..dc2e5257 100644 --- a/src/user/user_mesh.cc +++ b/src/user/user_mesh.cc @@ -4271,12 +4271,8 @@ void mjCFlex::Compile(const mjVFS* vfs) { } // linear elasticity - stiffness.assign(21*nelem, 0); - if (interpolated) { - int min_size = ceil(nodexpos.size()*nodexpos.size() / 21); - if (min_size > nelem) { - throw mjCError(this, "Trilinear dofs are require at least %d elements", "", min_size); - } + if (!interpolated) { + stiffness.assign(21 * nelem, 0); } // geometrically nonlinear elasticity @@ -4333,6 +4329,11 @@ void mjCFlex::Compile(const mjVFS* vfs) { } if (!stiffness_cached && young > 0 && interpolated) { + int n = pow(order_ + 1, 3); + int ndof = 3 * n; + if (stiffness.size() < ndof * ndof) { + stiffness.resize(ndof * ndof, 0); + } ComputeLinearStiffness(stiffness, nodexpos.data(), young, poisson, order_); } diff --git a/src/user/user_model.cc b/src/user/user_model.cc index 2af21381..f3206d13 100644 --- a/src/user/user_model.cc +++ b/src/user/user_model.cc @@ -1194,6 +1194,7 @@ void mjCModel::Clear() { nflexedge = 0; nflexelem = 0; nflexelemdata = 0; + nflexstiffness = 0; nflexelemedge = 0; nflexshelldata = 0; nflexevpair = 0; @@ -2177,6 +2178,7 @@ void mjCModel::SetSizes() { } nbvh = nbvhstatic + nbvhdynamic; + int extra_stiffness_size = 0; // flex counts for (int i=0; i < nflex; i++) { nflexnode += flexes_[i]->nnode; @@ -2188,6 +2190,9 @@ void mjCModel::SetSizes() { nflexshelldata += (int)flexes_[i]->shell.size(); nflexevpair += (int)flexes_[i]->evpair.size()/2; nflextexcoord += (flexes_[i]->HasTexcoord() ? flexes_[i]->get_texcoord().size()/2 : 0); + if (flexes_[i]->order_ != 0) { + extra_stiffness_size += (3 * flexes_[i]->nnode) * (3 * flexes_[i]->nnode); + } if (flexes_[i]->interpolated || flexes_[i]->rigid) { continue; } @@ -2237,6 +2242,9 @@ void mjCModel::SetSizes() { } } } + // TODO: This can be compacted further when we update mjwarp to not rely on + // 21*elem_adr for non-interpolated flexes. + nflexstiffness = nflexelem * 21 + extra_stiffness_size; // mesh counts for (int i=0; i < nmesh; i++) { @@ -3435,6 +3443,8 @@ void mjCModel::CopyObjects(mjModel* m) { shelldata_adr = 0; evpair_adr = 0; texcoord_adr = 0; + int standard_stiffness_size = 21 * m->nflexelem; + int current_extra_stiffness_adr = standard_stiffness_size; for (int i=0; i < nflex; i++) { // get pointer mjCFlex* pfl = flexes_[i]; @@ -3457,10 +3467,18 @@ void mjCModel::CopyObjects(mjModel* m) { mjuu_copyvec(m->flex_rgba + 4 * i, pfl->rgba, 4); // elasticity - if (!pfl->stiffness.empty()) { - mjuu_copyvec(m->flex_stiffness + 21 * elem_adr, pfl->stiffness.data(), pfl->stiffness.size()); + if (pfl->order_ == 0) { + m->flex_stiffnessadr[i] = 21 * elem_adr; } else { - mjuu_zerovec(m->flex_stiffness + 21 * elem_adr, 21 * pfl->nelem); + m->flex_stiffnessadr[i] = current_extra_stiffness_adr; + current_extra_stiffness_adr += (3 * pfl->nnode) * (3 * pfl->nnode); + } + + if (!pfl->stiffness.empty()) { + mjuu_copyvec(m->flex_stiffness + m->flex_stiffnessadr[i], pfl->stiffness.data(), pfl->stiffness.size()); + } else { + int size = (pfl->order_ == 0) ? 21 * pfl->nelem : (3 * pfl->nnode) * (3 * pfl->nnode); + mjuu_zerovec(m->flex_stiffness + m->flex_stiffnessadr[i], size); } if (!pfl->bending.empty()) { mjuu_copyvec(m->flex_bending + 17 * edge_adr, pfl->bending.data(), pfl->bending.size()); @@ -5125,12 +5143,12 @@ void mjCModel::TryCompile(mjModel*& m, mjData*& d, const mjVFS* vfs) { mj_makeModel(&m, nq, nv, nu, na, nbody, nbvh, nbvhstatic, nbvhdynamic, noct, njnt, ntree, nM, nB, nC, nD, ngeom, nsite, ncam, nlight, nflex, nflexnode, nflexvert, nflexedge, nflexelem, - nflexelemdata, nflexelemedge, nflexshelldata, nflexevpair, nflextexcoord, nJfe, nJfv, - nmesh, nmeshvert, nmeshnormal, nmeshtexcoord, nmeshface, nmeshgraph, nmeshpoly, - nmeshpolyvert, nmeshpolymap, nskin, nskinvert, nskintexvert, nskinface, nskinbone, - nskinbonevert, nhfield, nhfielddata, ntex, ntexdata, nmat, npair, nexclude, - neq, ntendon, nJten, nwrap, nsensor, nnumeric, nnumericdata, ntext, ntextdata, - ntuple, ntupledata, nkey, nmocap, nplugin, npluginattr, + nflexelemdata, nflexstiffness, nflexelemedge, nflexshelldata, nflexevpair, + nflextexcoord, nJfe, nJfv, nmesh, nmeshvert, nmeshnormal, nmeshtexcoord, nmeshface, + nmeshgraph, nmeshpoly, nmeshpolyvert, nmeshpolymap, nskin, nskinvert, nskintexvert, + nskinface, nskinbone, nskinbonevert, nhfield, nhfielddata, ntex, ntexdata, nmat, + npair, nexclude, neq, ntendon, nJten, nwrap, nsensor, nnumeric, nnumericdata, ntext, + ntextdata, ntuple, ntupledata, nkey, nmocap, nplugin, npluginattr, nuser_body, nuser_jnt, nuser_geom, nuser_site, nuser_cam, nuser_tendon, nuser_actuator, nuser_sensor, nnames, npaths); if (!m) { diff --git a/src/user/user_model.h b/src/user/user_model.h index 90feb1ed..9513b27b 100644 --- a/src/user/user_model.h +++ b/src/user/user_model.h @@ -96,6 +96,7 @@ class mjCModel_ : public mjsElement { mjtSize nflexedge; // number of edges in all flexes mjtSize nflexelem; // number of elements in all flexes mjtSize nflexelemdata; // number of element vertex ids in all flexes + mjtSize nflexstiffness; // number of stiffness parameters in all flexes mjtSize nflexelemedge; // number of element edges in all flexes mjtSize nflexshelldata; // number of shell fragment vertex ids in all flexes mjtSize nflexevpair; // number of element-vertex pairs in all flexes diff --git a/test/user/user_flex_test.cc b/test/user/user_flex_test.cc index 017bb320..478b23e0 100644 --- a/test/user/user_flex_test.cc +++ b/test/user/user_flex_test.cc @@ -447,7 +447,7 @@ TEST_F(UserFlexTest, StiffnessMatrix) { std::array error; mjModel* m = LoadModelFromString(xml, error.data(), error.size()); ASSERT_THAT(m, NotNull()) << error.data(); - EXPECT_NE(m->flex_stiffness[0], 0); + EXPECT_NE(m->flex_stiffness[m->flex_stiffnessadr[0]], 0); EXPECT_EQ(m->nflexnode, 8); // constants are in the kernel @@ -456,7 +456,8 @@ TEST_F(UserFlexTest, StiffnessMatrix) { zeros[i] = 0; ones[i] = 1; } - mju_mulMatVec(res, m->flex_stiffness, ones, 3*m->nflexnode, 3*m->nflexnode); + mju_mulMatVec(res, m->flex_stiffness + m->flex_stiffnessadr[0], ones, + 3 * m->nflexnode, 3 * m->nflexnode); EXPECT_THAT(res, Pointwise(MjNear(1e-8, 1e-4), zeros)); mj_deleteModel(m); @@ -496,7 +497,8 @@ TEST_F(UserFlexTest, StiffnessCacheDiffersByGeometry) { // Same number of nodes but different stiffness due to different geometry EXPECT_EQ(m_small->nflexnode, m_large->nflexnode); - EXPECT_NE(m_small->flex_stiffness[0], m_large->flex_stiffness[0]); + EXPECT_NE(m_small->flex_stiffness[m_small->flex_stiffnessadr[0]], + m_large->flex_stiffness[m_large->flex_stiffnessadr[0]]); mj_deleteModel(m_small); mj_deleteModel(m_large); diff --git a/unity/Runtime/Bindings/MjBindings.cs b/unity/Runtime/Bindings/MjBindings.cs index ca287bd6..72db845d 100644 --- a/unity/Runtime/Bindings/MjBindings.cs +++ b/unity/Runtime/Bindings/MjBindings.cs @@ -962,6 +962,7 @@ public unsafe struct mjModel_ { public UInt64 nflexedge; public UInt64 nflexelem; public UInt64 nflexelemdata; + public UInt64 nflexstiffness; public UInt64 nflexelemedge; public UInt64 nflexshelldata; public UInt64 nflexevpair; @@ -1208,6 +1209,7 @@ public unsafe struct mjModel_ { public int* flex_elemadr; public int* flex_elemnum; public int* flex_elemdataadr; + public int* flex_stiffnessadr; public int* flex_elemedgeadr; public int* flex_shellnum; public int* flex_shelldataadr; diff --git a/wasm/codegen/generated/bindings.cc b/wasm/codegen/generated/bindings.cc index 410cef5a..3fcacd7d 100644 --- a/wasm/codegen/generated/bindings.cc +++ b/wasm/codegen/generated/bindings.cc @@ -3737,6 +3737,12 @@ struct MjModel { void set_nflexelemdata(int value) { ptr_->nflexelemdata = static_cast(value); } + int nflexstiffness() const { + return static_cast(ptr_->nflexstiffness); + } + void set_nflexstiffness(int value) { + ptr_->nflexstiffness = static_cast(value); + } int nflexelemedge() const { return static_cast(ptr_->nflexelemedge); } @@ -4661,6 +4667,9 @@ struct MjModel { emscripten::val flex_elemdataadr() const { return emscripten::val(emscripten::typed_memory_view(ptr_->nflex, ptr_->flex_elemdataadr)); } + emscripten::val flex_stiffnessadr() const { + return emscripten::val(emscripten::typed_memory_view(ptr_->nflex, ptr_->flex_stiffnessadr)); + } emscripten::val flex_elemedgeadr() const { return emscripten::val(emscripten::typed_memory_view(ptr_->nflex, ptr_->flex_elemedgeadr)); } @@ -4746,7 +4755,7 @@ struct MjModel { return emscripten::val(emscripten::typed_memory_view(ptr_->nflex * 3, ptr_->flex_size)); } emscripten::val flex_stiffness() const { - return emscripten::val(emscripten::typed_memory_view(ptr_->nflexelem * 21, ptr_->flex_stiffness)); + return emscripten::val(emscripten::typed_memory_view(ptr_->nflexstiffness, ptr_->flex_stiffness)); } emscripten::val flex_bending() const { return emscripten::val(emscripten::typed_memory_view(ptr_->nflexedge * 17, ptr_->flex_bending)); @@ -11848,6 +11857,7 @@ EMSCRIPTEN_BINDINGS(mujoco_bindings) { .property("flex_solmix", &MjModel::flex_solmix) .property("flex_solref", &MjModel::flex_solref) .property("flex_stiffness", &MjModel::flex_stiffness) + .property("flex_stiffnessadr", &MjModel::flex_stiffnessadr) .property("flex_texcoord", &MjModel::flex_texcoord) .property("flex_texcoordadr", &MjModel::flex_texcoordadr) .property("flex_vert", &MjModel::flex_vert) @@ -12045,6 +12055,7 @@ EMSCRIPTEN_BINDINGS(mujoco_bindings) { .property("nflexevpair", &MjModel::nflexevpair, &MjModel::set_nflexevpair, reference()) .property("nflexnode", &MjModel::nflexnode, &MjModel::set_nflexnode, reference()) .property("nflexshelldata", &MjModel::nflexshelldata, &MjModel::set_nflexshelldata, reference()) + .property("nflexstiffness", &MjModel::nflexstiffness, &MjModel::set_nflexstiffness, reference()) .property("nflextexcoord", &MjModel::nflextexcoord, &MjModel::set_nflextexcoord, reference()) .property("nflexvert", &MjModel::nflexvert, &MjModel::set_nflexvert, reference()) .property("ngeom", &MjModel::ngeom, &MjModel::set_ngeom, reference())