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
This commit is contained in:
Alessio Quaglino
2026-04-14 09:38:17 -07:00
committed by Copybara-Service
parent 08bb6e66f1
commit 5724158593
14 changed files with 100 additions and 48 deletions
+1 -1
View File
@@ -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) {
+20 -22
View File
@@ -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) {
+1 -1
View File
@@ -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,
+1 -1
View File
@@ -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];
+7 -6
View File
@@ -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_);
}
+27 -9
View File
@@ -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) {
+1
View File
@@ -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