From dbe18f574006adbc410f1c0551a53a24e6156c6e Mon Sep 17 00:00:00 2001 From: Saran Tunyasuvunakool Date: Wed, 17 Jul 2024 10:03:53 -0700 Subject: [PATCH] Add additional MuJoCo fields to MJX structs. PiperOrigin-RevId: 653270676 Change-Id: I5f7a62d3bd542abc0d618e63d2426663e7d6406f --- doc/changelog.rst | 8 + mjx/mujoco/mjx/_src/io.py | 190 ++++++++----- mjx/mujoco/mjx/_src/smooth.py | 3 + mjx/mujoco/mjx/_src/types.py | 510 +++++++++++++++++++++++++++++----- 4 files changed, 569 insertions(+), 142 deletions(-) diff --git a/doc/changelog.rst b/doc/changelog.rst index 8b2080c2..6173007f 100644 --- a/doc/changelog.rst +++ b/doc/changelog.rst @@ -2,6 +2,14 @@ Changelog ========= +Upcoming version (not yet released) +----------------------------------- + +MJX +^^^ + +1. Added more fields to ``mjx.Model`` and ``mjx.Data`` for further compatibility with the corresponding MuJoCo structs. + Version 3.2.0 (Jul 15, 2024) ---------------------------- diff --git a/mjx/mujoco/mjx/_src/io.py b/mjx/mujoco/mjx/_src/io.py index efa1da86..1cdb015f 100644 --- a/mjx/mujoco/mjx/_src/io.py +++ b/mjx/mujoco/mjx/_src/io.py @@ -15,7 +15,7 @@ """Functions to initialize, load, or save data.""" import copy -from typing import List, Tuple, Union +from typing import Any, Dict, List, Tuple, Union import jax from jax import numpy as jp @@ -110,8 +110,6 @@ def put_model(m: mujoco.MjModel, device=None) -> types.Model: mj_field_names = {f.name for f in types.Model.fields()} - mjx_only fields = {f: getattr(m, f) for f in mj_field_names} fields['geom_rbound_hfield'] = fields['geom_rbound'] - fields['geom_rgba'] = fields['geom_rgba'].reshape((-1, 4)) - fields['mat_rgba'] = fields['mat_rgba'].reshape((-1, 4)) fields['cam_mat0'] = fields['cam_mat0'].reshape((-1, 3, 3)) fields['opt'] = _make_option(m.opt) fields['stat'] = _make_statistic(m.stat) @@ -137,33 +135,19 @@ def make_data(m: Union[types.Model, mujoco.MjModel]) -> types.Data: ne, nf, nl, nc = constraint.counts(efc_type) ncon, nefc = dim.size, ne + nf + nl + nc - zero_0 = jp.zeros(0, dtype=float) - zero_nv = jp.zeros(m.nv, dtype=float) - zero_nv_6 = jp.zeros((m.nv, 6), dtype=float) - zero_nv_nv = jp.zeros((m.nv, m.nv), dtype=float) - zero_nbody_3 = jp.zeros((m.nbody, 3), dtype=float) - zero_nbody_6 = jp.zeros((m.nbody, 6), dtype=float) - zero_nbody_10 = jp.zeros((m.nbody, 10), dtype=float) - zero_nbody_3_3 = jp.zeros((m.nbody, 3, 3), dtype=float) - zero_nefc = jp.zeros(nefc, dtype=float) - zero_na = jp.zeros(m.na, dtype=float) - zero_nu = jp.zeros(m.nu, dtype=float) - zero_njnt_3 = jp.zeros((m.njnt, 3), dtype=float) - zero_nm = jp.zeros(m.nM, dtype=float) - contact = types.Contact( - dist=jp.zeros(ncon), - pos=jp.zeros((ncon, 3)), - frame=jp.zeros((ncon, 3, 3)), - includemargin=jp.zeros(ncon), - friction=jp.zeros((ncon, 5)), - solref=jp.zeros((ncon, mujoco.mjNREF)), - solreffriction=jp.zeros((ncon, mujoco.mjNREF)), - solimp=jp.zeros((ncon, mujoco.mjNIMP)), + dist=jp.zeros((ncon,), dtype=float), + pos=jp.zeros((ncon, 3), dtype=float), + frame=jp.zeros((ncon, 3, 3), dtype=float), + includemargin=jp.zeros((ncon,), dtype=float), + friction=jp.zeros((ncon, 5), dtype=float), + solref=jp.zeros((ncon, mujoco.mjNREF), dtype=float), + solreffriction=jp.zeros((ncon, mujoco.mjNREF), dtype=float), + solimp=jp.zeros((ncon, mujoco.mjNIMP), dtype=float), dim=dim, - geom1=jp.zeros(ncon, dtype=int) - 1, - geom2=jp.zeros(ncon, dtype=int) - 1, - geom=jp.zeros((ncon, 2), dtype=int) - 1, + geom1=jp.full((ncon,), -1, dtype=int), + geom2=jp.full((ncon,), -1, dtype=int), + geom=jp.full((ncon, 2), -1, dtype=int), efc_address=efc_address, ) @@ -173,59 +157,118 @@ def make_data(m: Union[types.Model, mujoco.MjModel]) -> types.Data: nl=nl, nefc=nefc, ncon=ncon, - solver_niter=jp.array(0, dtype=int), - time=jp.array(0.0, dtype=float), + solver_niter=jp.zeros((), dtype=int), + time=jp.zeros((), dtype=float), qpos=jp.array(m.qpos0), - qvel=zero_nv, - act=zero_na, - qacc_warmstart=zero_nv, - ctrl=zero_nu, - qfrc_applied=zero_nv, - xfrc_applied=zero_nbody_6, - eq_active=jp.zeros(m.neq, dtype=jp.uint8), - qacc=zero_nv, - act_dot=zero_na, - xpos=zero_nbody_3, + qvel=jp.zeros((m.nv,), dtype=float), + act=jp.zeros((m.na,), dtype=float), + qacc_warmstart=jp.zeros((m.nv,), dtype=float), + ctrl=jp.zeros((m.nu,), dtype=float), + qfrc_applied=jp.zeros((m.nv,), dtype=float), + xfrc_applied=jp.zeros((m.nbody, 6), dtype=float), + eq_active=jp.zeros((m.neq,), dtype=jp.uint8), + mocap_pos=jp.zeros((m.nmocap, 3), dtype=float), + mocap_quat=jp.zeros((m.nmocap, 4), dtype=float), + qacc=jp.zeros((m.nv,), dtype=float), + act_dot=jp.zeros((m.na,), dtype=float), + userdata=jp.zeros((m.nuserdata,), dtype=float), + sensordata=jp.zeros((m.nsensordata,), dtype=float), + xpos=jp.zeros((m.nbody, 3), dtype=float), xquat=jp.zeros((m.nbody, 4), dtype=float), - xmat=zero_nbody_3_3, - xipos=zero_nbody_3, - ximat=zero_nbody_3_3, - xanchor=zero_njnt_3, - xaxis=zero_njnt_3, + xmat=jp.zeros((m.nbody, 3, 3), dtype=float), + xipos=jp.zeros((m.nbody, 3), dtype=float), + ximat=jp.zeros((m.nbody, 3, 3), dtype=float), + xanchor=jp.zeros((m.njnt, 3), dtype=float), + xaxis=jp.zeros((m.njnt, 3), dtype=float), geom_xpos=jp.zeros((m.ngeom, 3), dtype=float), geom_xmat=jp.zeros((m.ngeom, 3, 3), dtype=float), site_xpos=jp.zeros((m.nsite, 3), dtype=float), site_xmat=jp.zeros((m.nsite, 3, 3), dtype=float), cam_xpos=jp.zeros((m.ncam, 3), dtype=float), cam_xmat=jp.zeros((m.ncam, 3, 3), dtype=float), - subtree_com=zero_nbody_3, - cdof=zero_nv_6, - cinert=zero_nbody_10, - actuator_length=zero_nu, + light_xpos=jp.zeros((m.nlight, 3), dtype=float), + light_xdir=jp.zeros((m.nlight, 3), dtype=float), + subtree_com=jp.zeros((m.nbody, 3), dtype=float), + cdof=jp.zeros((m.nv, 6), dtype=float), + cinert=jp.zeros((m.nbody, 10), dtype=float), + flexvert_xpos=jp.zeros((m.nflexvert, 3), dtype=float), + flexelem_aabb=jp.zeros((m.nflexelem, 6), dtype=float), + flexedge_J_rownnz=jp.zeros((m.nflexedge,), dtype=int), + flexedge_J_rowadr=jp.zeros((m.nflexedge,), dtype=int), + flexedge_J_colind=jp.zeros((m.nflexedge, m.nv), dtype=int), + flexedge_J=jp.zeros((m.nflexedge, m.nv), dtype=float), + flexedge_length=jp.zeros((m.nflexedge,), dtype=float), + ten_wrapadr=jp.zeros((m.ntendon,), dtype=int), + ten_wrapnum=jp.zeros((m.ntendon,), dtype=int), + ten_J_rownnz=jp.zeros((m.ntendon,), dtype=int), + ten_J_rowadr=jp.zeros((m.ntendon,), dtype=int), + ten_J_colind=jp.zeros((m.ntendon, m.nv), dtype=int), + ten_J=jp.zeros((m.ntendon, m.nv), dtype=float), + ten_length=jp.zeros((m.ntendon,), dtype=float), + wrap_obj=jp.zeros((m.nwrap, 2), dtype=int), + wrap_xpos=jp.zeros((m.nwrap, 6), dtype=float), + actuator_length=jp.zeros((m.nu,), dtype=float), actuator_moment=jp.zeros((m.nu, m.nv), dtype=float), - crb=zero_nbody_10, - qM=zero_nm if support.is_sparse(m) else zero_nv_nv, - qLD=zero_nm if support.is_sparse(m) else zero_nv_nv, - qLDiagInv=zero_nv if support.is_sparse(m) else zero_0, + crb=jp.zeros((m.nbody, 10), dtype=float), + qM=( + jp.zeros((m.nM,), dtype=float) + if support.is_sparse(m) + else jp.zeros((m.nv, m.nv), dtype=float) + ), + qLD=( + jp.zeros((m.nM,), dtype=float) + if support.is_sparse(m) + else jp.zeros((m.nv, m.nv), dtype=float) + ), + qLDiagInv=( + jp.zeros((m.nv,), dtype=float) if support.is_sparse(m) + else jp.zeros((0,), dtype=float) + ), + qLDiagSqrtInv=jp.zeros((m.nv,), dtype=float), + bvh_aabb_dyn=jp.zeros((m.nbvhdynamic, 6), dtype=float), + bvh_active=jp.zeros((m.nbvh,), dtype=jp.uint8), + flexedge_velocity=jp.zeros((m.nflexedge,), dtype=float), + ten_velocity=jp.zeros((m.ntendon,), dtype=float), + actuator_velocity=jp.zeros((m.nu,), dtype=float), + cvel=jp.zeros((m.nbody, 6), dtype=float), + cdof_dot=jp.zeros((m.nv, 6), dtype=float), + qfrc_bias=jp.zeros((m.nv,), dtype=float), + qfrc_spring=jp.zeros((m.nv,), dtype=float), + qfrc_damper=jp.zeros((m.nv,), dtype=float), + qfrc_gravcomp=jp.zeros((m.nv,), dtype=float), + qfrc_fluid=jp.zeros((m.nv,), dtype=float), + qfrc_passive=jp.zeros((m.nv,), dtype=float), + subtree_linvel=jp.zeros((m.nbody, 3), dtype=float), + subtree_angmom=jp.zeros((m.nbody, 3), dtype=float), + qH=jp.zeros((m.nM,), dtype=float), + qHDiagInv=jp.zeros((m.nv,), dtype=float), + D_rownnz=jp.zeros((m.nv,), dtype=int), + D_rowadr=jp.zeros((m.nv,), dtype=int), + D_colind=jp.zeros((m.nD,), dtype=int), + B_rownnz=jp.zeros((m.nbody,), dtype=int), + B_rowadr=jp.zeros((m.nbody,), dtype=int), + B_colind=jp.zeros((m.nB,), dtype=int), + qDeriv=jp.zeros((m.nD,), dtype=float), + qLU=jp.zeros((m.nD,), dtype=float), + actuator_force=jp.zeros((m.nu,), dtype=float), + qfrc_actuator=jp.zeros((m.nv,), dtype=float), + qfrc_smooth=jp.zeros((m.nv,), dtype=float), + qacc_smooth=jp.zeros((m.nv,), dtype=float), + qfrc_constraint=jp.zeros((m.nv,), dtype=float), + qfrc_inverse=jp.zeros((m.nv,), dtype=float), + cacc=jp.zeros((m.nbody, 6), dtype=float), + cfrc_int=jp.zeros((m.nbody, 6), dtype=float), + cfrc_ext=jp.zeros((m.nbody, 6), dtype=float), contact=contact, efc_type=efc_type, efc_J=jp.zeros((nefc, m.nv), dtype=float), - efc_frictionloss=zero_nefc, - efc_D=zero_nefc, - actuator_velocity=zero_nu, - cvel=zero_nbody_6, - cdof_dot=zero_nv_6, - qfrc_bias=zero_nv, - qfrc_gravcomp=zero_nv, - qfrc_passive=zero_nv, - efc_aref=zero_nefc, - qfrc_actuator=zero_nv, - qfrc_smooth=zero_nv, - qacc_smooth=zero_nv, - qfrc_constraint=zero_nv, - qfrc_inverse=zero_nv, - efc_force=zero_nefc, - userdata=jp.zeros(m.nuserdata, dtype=float), + efc_frictionloss=jp.zeros((nefc,), dtype=float), + efc_D=jp.zeros((nefc,), dtype=float), + efc_aref=jp.zeros((nefc,), dtype=float), + efc_force=jp.zeros((nefc,), dtype=float), + _qM_sparse=jp.zeros((m.nM), dtype=float), + _qLD_sparse=jp.zeros((m.nM), dtype=float), + _qLDiagInv_sparse=jp.zeros((m.nv,), dtype=float), ) return d @@ -296,6 +339,9 @@ def get_data_into( result_i.efc_J_colind[:] = np.tile(np.arange(m.nv), nefc) for field in types.Data.fields(): + if field.name.startswith('_') and field.name.endswith('_sparse'): + continue + if field.name == 'contact': _get_contact(result_i.contact, d_i.contact) # efc_address must be updated because rows were deleted above: @@ -379,7 +425,8 @@ def put_data(m: mujoco.MjModel, d: mujoco.MjData, device=None) -> types.Data: if d_val > val: raise ValueError(f'd.{name} too high, d.{name} = {d_val}, model = {val}') - fields = {f.name: getattr(d, f.name) for f in types.Data.fields()} + fields = {f.name: getattr(d, f.name) for f in types.Data.fields() + if not f.name.endswith('_sparse')} # MJX prefers square matrices for these fields: for fname in ('xmat', 'ximat', 'geom_xmat', 'site_xmat', 'cam_xmat'): @@ -427,6 +474,9 @@ def put_data(m: mujoco.MjModel, d: mujoco.MjData, device=None) -> types.Data: fields[fname] = value # convert qM and qLD if jacobian is dense + fields['_qM_sparse'] = fields['qM'] + fields['_qLD_sparse'] = fields['qLD'] + fields['_qLDiagInv_sparse'] = fields['qLDiagInv'] if not support.is_sparse(m): fields['qM'] = np.zeros((m.nv, m.nv)) mujoco.mj_fullM(m, fields['qM'], d.qM) diff --git a/mjx/mujoco/mjx/_src/smooth.py b/mjx/mujoco/mjx/_src/smooth.py index a3fe0019..7f105199 100644 --- a/mjx/mujoco/mjx/_src/smooth.py +++ b/mjx/mujoco/mjx/_src/smooth.py @@ -282,6 +282,8 @@ def crb(m: Model, d: Data) -> Data: crb_cdof = jax.vmap(math.inert_mul)(crb_dof, d.cdof) qm = support.make_m(m, crb_cdof, d.cdof, m.dof_armature) d = d.replace(qM=qm) + if support.is_sparse(m): + d = d.replace(_qM_sparse=qm) return d @@ -340,6 +342,7 @@ def factor_m(m: Model, d: Data) -> Data: qld = (qld / qld[jp.array(madr_ds)]).at[m.dof_Madr].set(qld_diag) d = d.replace(qLD=qld, qLDiagInv=1 / qld_diag) + d = d.replace(_qLD_sparse=d.qLD, _qLDiagInv_sparse=d.qLDiagInv) return d diff --git a/mjx/mujoco/mjx/_src/types.py b/mjx/mujoco/mjx/_src/types.py index c85701d4..f20f1cbb 100644 --- a/mjx/mujoco/mjx/_src/types.py +++ b/mjx/mujoco/mjx/_src/types.py @@ -334,25 +334,49 @@ class Model(PyTreeNode): nu: number of actuators/controls = dim(ctrl) na: number of activation states = dim(act) nbody: number of bodies + nbvh: number of total bounding volumes in all bodies + nbvhstatic: number of static bounding volumes (aabb stored in mjModel) + nbvhdynamic: number of dynamic bounding volumes (aabb stored in mjData) njnt: number of joints ngeom: number of geoms nsite: number of sites ncam: number of cameras + nlight: number of lights + nflex: number of flexes + nflexvert: number of vertices in all flexes + nflexedge: number of edges in all flexes + nflexelem: number of elements in all flexes + nflexelemdata: number of element vertex ids in all flexes + nflexshelldata: number of shell fragment vertex ids in all flexes + nflexevpair: number of element-vertex pairs in all flexes + nflextexcoord: number of vertices with texture coordinates nmesh: number of meshes nmeshvert: number of vertices in all meshes + nmeshnormal: number of normals in all meshes + nmeshtexcoord: number of texcoords in all meshes nmeshface: number of triangular faces in all meshes + nmeshgraph: number of ints in mesh auxiliary data nhfield: number of heightfields + nhfielddata: number of data points in all heightfields nmat: number of materials npair: number of predefined geom pairs nexclude: number of excluded geom pairs neq: number of equality constraints - ngravcomp: number of bodies with nonzero gravcomp + ntendon: number of tendons + nwrap: number of wrap objects in all tendon paths + nsensor: number of sensors nnumeric: number of numeric custom fields ntuple: number of tuple custom fields - nsensor: number of sensors nkey: number of keyframes - nuserdata: size of userdata array + nmocap: number of mocap bodies nM: number of non-zeros in sparse inertia matrix + nD: number of non-zeros in sparse dof-dof matrix + nB: number of non-zeros in sparse body-dof matrix + ntree: number of kinematic trees under world body + ngravcomp: number of bodies with nonzero gravcomp + nuserdata: size of userdata array + nsensordata: number of mjtNums in sensor data vector + narena: number of bytes in the mjData arena (inclusive of stack) opt: physics options stat: model statistics qpos0: qpos values at default pose (nq,) @@ -364,8 +388,10 @@ class Model(PyTreeNode): body_jntadr: start addr of joints; -1: no joints (nbody,) body_dofnum: number of motion degrees of freedom (nbody,) body_dofadr: start addr of dofs; -1: no dofs (nbody,) + body_treeid: id of body's kinematic tree; -1: static (nbody,) body_geomnum: number of geoms (nbody,) body_geomadr: start addr of geoms; -1: no geoms (nbody,) + body_simple: 1: diag M; 2: diag M, sliders only (nbody,) body_pos: position offset rel. to parent body (nbody, 3) body_quat: orientation offset rel. to parent body (nbody, 4) body_ipos: local position of center of mass (nbody, 3) @@ -374,6 +400,14 @@ class Model(PyTreeNode): body_subtreemass: mass of subtree starting at this body (nbody,) body_inertia: diagonal inertia in ipos/iquat frame (nbody, 3) body_gravcomp: antigravity force, units of body weight (nbody,) + body_margin: MAX over all geom margins (nbody,) + body_contype: OR over all geom contypes (nbody,) + body_conaffinity: OR over all geom conaffinities (nbody,) + body_bvhadr: address of bvh root (nbody,) + body_bvhnum: number of bounding volumes (nbody,) + bvh_child: left and right children in tree (nbvh, 2) + bvh_nodeid: geom or elem id of node; -1: non-leaf (nbvh,) + bvh_aabb: local bounding box (center, size) (nbvhstatic, 6) body_invweight0: mean inv inert in qpos0 (trn, rot) (nbody, 2) jnt_type: type of joint (mjtJoint) (njnt,) jnt_qposadr: start addr in 'qpos' for joint's data (njnt,) @@ -394,7 +428,9 @@ class Model(PyTreeNode): dof_bodyid: id of dof's body (nv,) dof_jntid: id of dof's joint (nv,) dof_parentid: id of dof's parent; -1: none (nv,) + dof_treeid: id of dof's kinematic tree (nv,) dof_Madr: dof address in M-diagonal (nv,) + dof_simplenum: number of consecutive simple dofs (nv,) dof_solref: constraint solver reference:frictionloss (nv, mjNREF) dof_solimp: constraint solver impedance:frictionloss (nv, mjNIMP) dof_frictionloss: dof friction loss (nv,) @@ -415,6 +451,7 @@ class Model(PyTreeNode): geom_solref: constraint solver reference: contact (ngeom, mjNREF) geom_solimp: constraint solver impedance: contact (ngeom, mjNIMP) geom_size: geom-specific size parameters (ngeom, 3) + geom_aabb: bounding box, (center, size) (ngeom, 6) geom_rbound: radius of bounding sphere (ngeom,) geom_rbound_hfield: static rbound for hfield grid bounds (ngeom,) geom_pos: local position offset rel. to body (ngeom, 3) @@ -432,14 +469,71 @@ class Model(PyTreeNode): cam_pos: position rel. to body frame (ncam, 3) cam_quat: orientation rel. to body frame (ncam, 4) cam_poscom0: global position rel. to sub-com in qpos0 (ncam, 3) - cam_pos0: global position rel. to body in qpos0 (ncam, 3) - cam_mat0: global orientation in qpos0 (ncam, 9) + cam_pos0: global position rel. to body in qpos0 (ncam, 3) + cam_mat0: global orientation in qpos0 (ncam, 3, 3) + cam_fovy: y field-of-view (ncam,) + cam_resolution: resolution: pixels (ncam, 2) + cam_sensorsize: sensor size: length (ncam, 2) + cam_intrinsic: [focal length; principal point] (ncam, 4) + light_mode: light tracking mode (mjtCamLight) (nlight,) + light_bodyid: id of light's body (nlight,) + light_targetbodyid: id of targeted body; -1: none (nlight,) + light_pos: position rel. to body frame (nlight, 3) + light_dir: direction rel. to body frame (nlight, 3) + light_poscom0: global position rel. to sub-com in qpos0 (nlight, 3) + light_pos0: global position rel. to body in qpos0 (nlight, 3) + light_dir0: global direction in qpos0 (nlight, 3) + flex_contype: flex contact type (nflex,) + flex_conaffinity: flex contact affinity (nflex,) + flex_condim: contact dimensionality (1, 3, 4, 6) (nflex,) + flex_priority: flex contact priority (nflex,) + flex_solmix: mix coef for solref/imp in contact pair (nflex,) + flex_solref: constraint solver reference: contact (nflex, mjNREF) + flex_solimp: constraint solver impedance: contact (nflex, mjNIMP) + flex_friction: friction for (slide, spin, roll) (nflex,) + flex_margin: detect contact if dist