From 9ffd50ce6a76b5cea34375a746fe23ff6405805f Mon Sep 17 00:00:00 2001 From: Taylor Howell Date: Wed, 6 May 2026 08:47:18 -0700 Subject: [PATCH] Import google-deepmind/mujoco_warp from GitHub. PiperOrigin-RevId: 911361969 Change-Id: I2dd1223a4279b3df2f11adc4f44778b8d153e599 --- .../mjx/third_party/mujoco_warp/_src/cli.py | 15 +++---- .../mjx/third_party/mujoco_warp/_src/io.py | 41 +++++++++++++++++++ .../third_party/mujoco_warp/_src/passive.py | 22 ++++++---- .../third_party/mujoco_warp/_src/sensor.py | 17 +++++++- .../mjx/third_party/mujoco_warp/_src/types.py | 16 ++++++-- .../third_party/mujoco_warp/pyproject.toml | 1 + mjx/mujoco/mjx/warp/forward.py | 21 ++++++++-- mjx/mujoco/mjx/warp/render.py | 1 - mjx/mujoco/mjx/warp/types.py | 16 +++++++- 9 files changed, 122 insertions(+), 28 deletions(-) diff --git a/mjx/mujoco/mjx/third_party/mujoco_warp/_src/cli.py b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/cli.py index 912df58f..0e427408 100644 --- a/mjx/mujoco/mjx/third_party/mujoco_warp/_src/cli.py +++ b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/cli.py @@ -21,14 +21,12 @@ from typing import Callable, Tuple, get_type_hints import mujoco import numpy as np import warp as wp -from absl import app from absl import flags from etils import epath import mujoco.mjx.third_party.mujoco_warp as mjw from mujoco.mjx.third_party.mujoco_warp._src import warp_util -from mujoco.mjx.third_party.mujoco_warp._src.io import find_keys -from mujoco.mjx.third_party.mujoco_warp._src.io import make_trajectory +from mujoco.mjx.third_party.mujoco_warp._src.io import load_trajectory from mujoco.mjx.third_party.mujoco_warp._src.io import override_model from mujoco.mjx.third_party.mujoco_warp._src.util_misc import halton @@ -46,7 +44,7 @@ NOISE_STD = flags.DEFINE_float("noise_std", 0.01, "add noise to ctrl signal (sta NOISE_RATE = flags.DEFINE_float("noise_rate", 0.1, "add noise to ctrl signal (noise rate)") DEVICE = flags.DEFINE_string("device", None, "override the default Warp device") -REPLAY = flags.DEFINE_string("replay", None, "keyframe sequence to replay, keyframe name must prefix match") +REPLAY = flags.DEFINE_string("replay", None, "NPZ file with ctrl sequence to replay") RENDER_WIDTH = flags.DEFINE_integer("render_width", 64, "render width (pixels)") RENDER_HEIGHT = flags.DEFINE_integer("render_height", 64, "render height (pixels)") @@ -133,11 +131,10 @@ def init_structs( mjd = mujoco.MjData(mjm) ctrls = None if REPLAY.value: - keys = find_keys(mjm, REPLAY.value) - if not keys: - raise app.UsageError(f"Key prefix not found: {REPLAY.value}") - ctrls = make_trajectory(mjm, keys) - mujoco.mj_resetDataKeyframe(mjm, mjd, keys[0]) + ctrls = load_trajectory(REPLAY.value, mjm, mjd) + # default nstep to trajectory length when not explicitly set + if flags.FLAGS["nstep"].using_default_value: + flags.FLAGS.nstep = len(ctrls) elif mjm.nkey > 0 and KEYFRAME.value > -1: mujoco.mj_resetDataKeyframe(mjm, mjd, KEYFRAME.value) ctrls = [mjd.ctrl.copy() for _ in range(NSTEP.value)] diff --git a/mjx/mujoco/mjx/third_party/mujoco_warp/_src/io.py b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/io.py index dc834e06..309f9e41 100644 --- a/mjx/mujoco/mjx/third_party/mujoco_warp/_src/io.py +++ b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/io.py @@ -660,6 +660,14 @@ def put_model(mjm: mujoco.MjModel) -> types.Model: m.flexedge_J_rowadr = mjm.flexedge_J_rowadr m.flexedge_J_colind = mjm.flexedge_J_colind.reshape(-1) + # flex_bendingadr backward compat: flatten old (nflexedge, 17) to 1D + if not check_version("mujoco>=3.8.1.dev909088123"): + m.flex_bendingadr = ( + np.array([mjm.flex_edgeadr[i] * 17 for i in range(mjm.nflex)], dtype=int) if mjm.nflex else np.zeros(0, dtype=int) + ) + m.flex_bending = mjm.flex_bending.ravel() + m.nflexbending = len(m.flex_bending) + # place m on device sizes = dict({"*": 1}, **{f.name: getattr(m, f.name) for f in dataclasses.fields(types.Model) if f.type is int}) for f in dataclasses.fields(types.Model): @@ -2654,6 +2662,39 @@ def make_trajectory(model: mujoco.MjModel, keys: list[int]) -> np.ndarray: return np.array(ctrls) +def load_trajectory(npz_path: str, mjm: mujoco.MjModel, mjd: mujoco.MjData) -> np.ndarray: + """Load ctrl sequence from NPZ and interpolate to model timestep. + + If the trajectory dt differs from mjm.opt.timestep, each ctrl value is held + constant (zero-order hold) for the appropriate number of physics steps. + + The NPZ file should contain: + - 'ctrl': array of shape (nstep, nu) with ctrl values + - 'times': array of shape (nstep,) with timestamps + - 'qpos' (optional): array of shape (1, nq) - initial state + - 'qvel' (optional): array of shape (1, nv) - initial state + """ + data = np.load(npz_path) + ctrl = data["ctrl"] + times = data["times"] + + if ctrl.shape[1] != mjm.nu: + raise ValueError(f"ctrl shape {ctrl.shape} does not match model nu={mjm.nu}") + + # set initial state from first frame if available + if "qpos" in data and data["qpos"].shape[1] == mjm.nq: + mjd.qpos[:] = data["qpos"][0] + if "qvel" in data and data["qvel"].shape[1] == mjm.nv: + mjd.qvel[:] = data["qvel"][0] + + # determine decimation from timing + ctrl_dt = (times[1] - times[0]) if len(times) > 1 else mjm.opt.timestep + decimation = max(1, round(ctrl_dt / mjm.opt.timestep)) + + # expand: each ctrl held constant for decimation physics steps + return np.repeat(ctrl, decimation, axis=0) + + @wp.kernel def _build_rays( # In: diff --git a/mjx/mujoco/mjx/third_party/mujoco_warp/_src/passive.py b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/passive.py index 31469590..2d261abb 100644 --- a/mjx/mujoco/mjx/third_party/mujoco_warp/_src/passive.py +++ b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/passive.py @@ -604,12 +604,13 @@ def _flex_elasticity( flex_elemadr: wp.array[int], flex_elemnum: wp.array[int], flex_elemdataadr: wp.array[int], + flex_stiffnessadr: wp.array[int], flex_elemedgeadr: wp.array[int], flex_vertbodyid: wp.array[int], flex_elem: wp.array[int], flex_elemedge: wp.array[int], flexedge_length0: wp.array[float], - flex_stiffness: wp.array2d[float], + flex_stiffness: wp.array[float], flex_damping: wp.array[float], # Data in: flexvert_xpos_in: wp.array2d[wp.vec3], @@ -669,11 +670,13 @@ def _flex_elasticity( elongation[e] = deformed * deformed - reference * reference + (deformed * deformed - previous * previous) * kD metric = wp.matrix(0.0, shape=(6, 6)) + stiffness_size = nedge * (nedge + 1) / 2 + stiffness_adr = flex_stiffnessadr[f] + local_elemid * stiffness_size id = int(0) for ed1 in range(nedge): for ed2 in range(ed1, nedge): - metric[ed1, ed2] = flex_stiffness[elemid, id] - metric[ed2, ed1] = flex_stiffness[elemid, id] + metric[ed1, ed2] = flex_stiffness[stiffness_adr + id] + metric[ed2, ed1] = flex_stiffness[stiffness_adr + id] id += 1 force = wp.matrix(0.0, shape=(6, 3)) @@ -699,10 +702,11 @@ def _flex_bending( flex_vertadr: wp.array[int], flex_edgeadr: wp.array[int], flex_edgenum: wp.array[int], + flex_bendingadr: wp.array[int], flex_vertbodyid: wp.array[int], flex_edge: wp.array[wp.vec2i], flex_edgeflap: wp.array[wp.vec2i], - flex_bending: wp.array2d[float], + flex_bending: wp.array[float], # Data in: flexvert_xpos_in: wp.array2d[wp.vec3], # Data out: @@ -730,8 +734,10 @@ def _flex_bending( flex_vertadr[f] + flex_edgeflap[edgeid][1], ) + adr = flex_bendingadr[f] + frc = wp.matrix(0.0, shape=(4, 3)) - if flex_bending[edgeid, 16]: + if flex_bending[adr + 16]: v0 = flexvert_xpos_in[worldid, v[0]] v1 = flexvert_xpos_in[worldid, v[1]] v2 = flexvert_xpos_in[worldid, v[2]] @@ -746,8 +752,8 @@ def _flex_bending( for x in range(3): acc = float(0.0) for j in range(nvert): - acc += flex_bending[edgeid, 4 * i + j] * flexvert_xpos_in[worldid, v[j]][x] - force[i, x] = -(acc + flex_bending[edgeid, 16] * frc[i, x]) + acc += flex_bending[adr + 4 * i + j] * flexvert_xpos_in[worldid, v[j]][x] + force[i, x] = -(acc + flex_bending[adr + 16] * frc[i, x]) for i in range(nvert): bodyid = flex_vertbodyid[v[i]] @@ -827,6 +833,7 @@ def passive(m: Model, d: Data): m.flex_elemadr, m.flex_elemnum, m.flex_elemdataadr, + m.flex_stiffnessadr, m.flex_elemedgeadr, m.flex_vertbodyid, m.flex_elem, @@ -851,6 +858,7 @@ def passive(m: Model, d: Data): m.flex_vertadr, m.flex_edgeadr, m.flex_edgenum, + m.flex_bendingadr, m.flex_vertbodyid, m.flex_edge, m.flex_edgeflap, diff --git a/mjx/mujoco/mjx/third_party/mujoco_warp/_src/sensor.py b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/sensor.py index 381dfd66..0610612b 100644 --- a/mjx/mujoco/mjx/third_party/mujoco_warp/_src/sensor.py +++ b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/sensor.py @@ -464,6 +464,8 @@ def _sensor_pos( body_geomnum: wp.array[int], body_geomadr: wp.array[int], body_iquat: wp.array2d[wp.quat], + body_mass: wp.array2d[float], + body_subtreemass: wp.array2d[float], jnt_qposadr: wp.array[int], geom_type: wp.array[int], geom_bodyid: wp.array[int], @@ -684,7 +686,18 @@ def _sensor_pos( if objtype == ObjType.XBODY: xpos = xpos_in[worldid, objid] elif objtype == ObjType.BODY: - xpos = xipos_in[worldid, objid] + # for massless bodies with positive subtree mass (e.g., flex parents), + # xipos is the static body frame origin; use subtree_com instead + if objid > 0: + if ( + body_mass[worldid % body_mass.shape[0], objid] < MJ_MINVAL + and body_subtreemass[worldid % body_subtreemass.shape[0], objid] >= MJ_MINVAL + ): + xpos = subtree_com_in[worldid, objid] + else: + xpos = xipos_in[worldid, objid] + else: + xpos = xipos_in[worldid, objid] elif objtype == ObjType.GEOM: xpos = geom_xpos_in[worldid, objid] elif objtype == ObjType.SITE: @@ -830,6 +843,8 @@ def sensor_pos(m: Model, d: Data): m.body_geomnum, m.body_geomadr, m.body_iquat, + m.body_mass, + m.body_subtreemass, m.jnt_qposadr, m.geom_type, m.geom_bodyid, diff --git a/mjx/mujoco/mjx/third_party/mujoco_warp/_src/types.py b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/types.py index 2395a711..7a4c8242 100644 --- a/mjx/mujoco/mjx/third_party/mujoco_warp/_src/types.py +++ b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/types.py @@ -867,6 +867,8 @@ class Model: nflexedge: number of edges in all flexes nflexelem: number of elements in all flexes nflexelemdata: number of element vertex ids in all flexes + nflexstiffness: number of stiffness parameters in all flexes + nflexbending: number of bending parameters in all flexes nflexelemedge: number of element edge ids in all flexes nflexshelldata: number of shell fragment vertex ids in all flexes nJfe: number of non-zeros in sparse flexedge Jacobian @@ -1022,7 +1024,9 @@ class Model: flex_elemadr: first element address (nflex,) flex_elemnum: number of elements (nflex,) flex_elemdataadr: first element vertex id address (nflex,) + flex_stiffnessadr: stiffness matrix address (nflex,) flex_elemedgeadr: first element edge id address (nflex,) + flex_bendingadr: first bending data address (nflex,) flex_shellnum: number of shells (nflex,) flex_shelldataadr: first shell data address (nflex,) flex_vertbodyid: vertex body ids (nflexvert,) @@ -1035,8 +1039,8 @@ class Model: flexedge_length0: edge lengths in qpos0 (nflexedge,) flexedge_invweight0: inv. inertia for the edge (nflexedge,) flex_radius: radius around primitive element (nflex,) - flex_stiffness: finite element stiffness matrix (nflexelem, 21) - flex_bending: bending stiffness (nflexedge, 17) + flex_stiffness: finite element stiffness matrix (nflexstiffness,) + flex_bending: bending stiffness (nflexbending,) flex_damping: Rayleigh's damping coefficient (nflex,) flex_centered: flex vertices are centered at body origin (nflex,) flexedge_J_rownnz: number of nonzeros in Jacobian row (nflexedge,) @@ -1266,6 +1270,8 @@ class Model: nflexedge: int nflexelem: int nflexelemdata: int + nflexstiffness: int + nflexbending: int nflexelemedge: int nflexshelldata: int nJfe: int @@ -1421,7 +1427,9 @@ class Model: flex_elemadr: array("nflex", int) flex_elemnum: array("nflex", int) flex_elemdataadr: array("nflex", int) + flex_stiffnessadr: array("nflex", int) flex_elemedgeadr: array("nflex", int) + flex_bendingadr: array("nflex", int) flex_shellnum: array("nflex", int) flex_shelldataadr: array("nflex", int) flex_vertbodyid: array("nflexvert", int) @@ -1434,8 +1442,8 @@ class Model: flexedge_length0: array("nflexedge", float) flexedge_invweight0: array("nflexedge", float) flex_radius: array("nflex", float) - flex_stiffness: array("nflexelem", 21, float) - flex_bending: array("nflexedge", 17, float) + flex_stiffness: array("nflexstiffness", float) + flex_bending: array("nflexbending", float) flex_damping: array("nflex", float) flex_centered: array("nflex", bool) flexedge_J_rownnz: array("nflexedge", int) diff --git a/mjx/mujoco/mjx/third_party/mujoco_warp/pyproject.toml b/mjx/mujoco/mjx/third_party/mujoco_warp/pyproject.toml index a5caa4b4..92750398 100644 --- a/mjx/mujoco/mjx/third_party/mujoco_warp/pyproject.toml +++ b/mjx/mujoco/mjx/third_party/mujoco_warp/pyproject.toml @@ -22,6 +22,7 @@ classifiers = [ "Programming Language :: Python :: 3.11", "Programming Language :: Python :: 3.12", "Programming Language :: Python :: 3.13", + "Programming Language :: Python :: 3.14", "Topic :: Scientific/Engineering", ] requires-python = ">=3.10" diff --git a/mjx/mujoco/mjx/warp/forward.py b/mjx/mujoco/mjx/warp/forward.py index f260ea92..0adf9236 100644 --- a/mjx/mujoco/mjx/warp/forward.py +++ b/mjx/mujoco/mjx/warp/forward.py @@ -46,6 +46,7 @@ _cb = mjwp_types.Callback( **{f.name: None for f in dataclasses.fields(mjwp_types.Callback) if f.init} ) + @ffi.format_args_for_warp def _forward_shim( # Model @@ -138,7 +139,8 @@ def _forward_shim( eq_ten_adr: wp.array[int], eq_type: wp.array[int], eq_wld_adr: wp.array[int], - flex_bending: wp.array2d[float], + flex_bending: wp.array[float], + flex_bendingadr: wp.array[int], flex_centered: wp.array[bool], flex_conaffinity: wp.array[int], flex_condim: wp.array[int], @@ -166,7 +168,8 @@ def _forward_shim( flex_solimp: wp.array[mjwp_types.vec5], flex_solmix: wp.array[float], flex_solref: wp.array[wp.vec2], - flex_stiffness: wp.array2d[float], + flex_stiffness: wp.array[float], + flex_stiffnessadr: wp.array[int], flex_vert: wp.array[wp.vec3], flex_vertadr: wp.array[int], flex_vertbodyid: wp.array[int], @@ -624,6 +627,7 @@ def _forward_shim( _m.eq_type = eq_type _m.eq_wld_adr = eq_wld_adr _m.flex_bending = flex_bending + _m.flex_bendingadr = flex_bendingadr _m.flex_centered = flex_centered _m.flex_conaffinity = flex_conaffinity _m.flex_condim = flex_condim @@ -652,6 +656,7 @@ def _forward_shim( _m.flex_solmix = flex_solmix _m.flex_solref = flex_solref _m.flex_stiffness = flex_stiffness + _m.flex_stiffnessadr = flex_stiffnessadr _m.flex_vert = flex_vert _m.flex_vertadr = flex_vertadr _m.flex_vertbodyid = flex_vertbodyid @@ -1504,6 +1509,7 @@ def _forward_jax_impl(m: types.Model, d: types.Data): m.eq_type, m._impl.eq_wld_adr, m._impl.flex_bending, + m._impl.flex_bendingadr, m._impl.flex_centered, m._impl.flex_conaffinity, m._impl.flex_condim, @@ -1532,6 +1538,7 @@ def _forward_jax_impl(m: types.Model, d: types.Data): m._impl.flex_solmix, m._impl.flex_solref, m._impl.flex_stiffness, + m._impl.flex_stiffnessadr, m._impl.flex_vert, m.flex_vertadr, m._impl.flex_vertbodyid, @@ -2107,7 +2114,8 @@ def _step_shim( eq_ten_adr: wp.array[int], eq_type: wp.array[int], eq_wld_adr: wp.array[int], - flex_bending: wp.array2d[float], + flex_bending: wp.array[float], + flex_bendingadr: wp.array[int], flex_centered: wp.array[bool], flex_conaffinity: wp.array[int], flex_condim: wp.array[int], @@ -2135,7 +2143,8 @@ def _step_shim( flex_solimp: wp.array[mjwp_types.vec5], flex_solmix: wp.array[float], flex_solref: wp.array[wp.vec2], - flex_stiffness: wp.array2d[float], + flex_stiffness: wp.array[float], + flex_stiffnessadr: wp.array[int], flex_vert: wp.array[wp.vec3], flex_vertadr: wp.array[int], flex_vertbodyid: wp.array[int], @@ -2595,6 +2604,7 @@ def _step_shim( _m.eq_type = eq_type _m.eq_wld_adr = eq_wld_adr _m.flex_bending = flex_bending + _m.flex_bendingadr = flex_bendingadr _m.flex_centered = flex_centered _m.flex_conaffinity = flex_conaffinity _m.flex_condim = flex_condim @@ -2623,6 +2633,7 @@ def _step_shim( _m.flex_solmix = flex_solmix _m.flex_solref = flex_solref _m.flex_stiffness = flex_stiffness + _m.flex_stiffnessadr = flex_stiffnessadr _m.flex_vert = flex_vert _m.flex_vertadr = flex_vertadr _m.flex_vertbodyid = flex_vertbodyid @@ -3489,6 +3500,7 @@ def _step_jax_impl(m: types.Model, d: types.Data): m.eq_type, m._impl.eq_wld_adr, m._impl.flex_bending, + m._impl.flex_bendingadr, m._impl.flex_centered, m._impl.flex_conaffinity, m._impl.flex_condim, @@ -3517,6 +3529,7 @@ def _step_jax_impl(m: types.Model, d: types.Data): m._impl.flex_solmix, m._impl.flex_solref, m._impl.flex_stiffness, + m._impl.flex_stiffnessadr, m._impl.flex_vert, m.flex_vertadr, m._impl.flex_vertbodyid, diff --git a/mjx/mujoco/mjx/warp/render.py b/mjx/mujoco/mjx/warp/render.py index 719e8df9..cf6754a1 100644 --- a/mjx/mujoco/mjx/warp/render.py +++ b/mjx/mujoco/mjx/warp/render.py @@ -48,7 +48,6 @@ _cb = mjwp_types.Callback( **{f.name: None for f in dataclasses.fields(mjwp_types.Callback) if f.init} ) - @ffi.format_args_for_warp def _render_shim( # Model diff --git a/mjx/mujoco/mjx/warp/types.py b/mjx/mujoco/mjx/warp/types.py index 81e77ec8..f64b6fba 100644 --- a/mjx/mujoco/mjx/warp/types.py +++ b/mjx/mujoco/mjx/warp/types.py @@ -143,6 +143,7 @@ class ModelWarp(PyTreeNode): eq_ten_adr: np.ndarray eq_wld_adr: np.ndarray flex_bending: np.ndarray + flex_bendingadr: np.ndarray flex_centered: np.ndarray flex_conaffinity: np.ndarray flex_condim: np.ndarray @@ -171,6 +172,7 @@ class ModelWarp(PyTreeNode): flex_solmix: np.ndarray flex_solref: np.ndarray flex_stiffness: np.ndarray + flex_stiffnessadr: np.ndarray flex_vert: np.ndarray flex_vertbodyid: np.ndarray flex_vertflexid: np.ndarray @@ -205,11 +207,13 @@ class ModelWarp(PyTreeNode): nJfe: int nacttrnbody: int nbranch: int + nflexbending: int nflexedge: int nflexelem: int nflexelemdata: int nflexelemedge: int nflexshelldata: int + nflexstiffness: int nflexvert: int nmaxcondim: int nmaxmeshdeg: int @@ -645,7 +649,8 @@ _NDIM = { 'eq_type': 1, 'eq_wld_adr': 1, 'exclude_signature': 1, - 'flex_bending': 2, + 'flex_bending': 1, + 'flex_bendingadr': 1, 'flex_centered': 1, 'flex_conaffinity': 1, 'flex_condim': 1, @@ -673,7 +678,8 @@ _NDIM = { 'flex_solimp': 2, 'flex_solmix': 1, 'flex_solref': 2, - 'flex_stiffness': 2, + 'flex_stiffness': 1, + 'flex_stiffnessadr': 1, 'flex_vert': 2, 'flex_vertadr': 1, 'flex_vertbodyid': 1, @@ -785,11 +791,13 @@ _NDIM = { 'neq': 0, 'nexclude': 0, 'nflex': 0, + 'nflexbending': 0, 'nflexedge': 0, 'nflexelem': 0, 'nflexelemdata': 0, 'nflexelemedge': 0, 'nflexshelldata': 0, + 'nflexstiffness': 0, 'nflexvert': 0, 'ngeom': 0, 'ngravcomp': 0, @@ -1224,6 +1232,7 @@ _BATCH_DIM = { 'eq_wld_adr': False, 'exclude_signature': False, 'flex_bending': False, + 'flex_bendingadr': False, 'flex_centered': False, 'flex_conaffinity': False, 'flex_condim': False, @@ -1252,6 +1261,7 @@ _BATCH_DIM = { 'flex_solmix': False, 'flex_solref': False, 'flex_stiffness': False, + 'flex_stiffnessadr': False, 'flex_vert': False, 'flex_vertadr': False, 'flex_vertbodyid': False, @@ -1363,11 +1373,13 @@ _BATCH_DIM = { 'neq': False, 'nexclude': False, 'nflex': False, + 'nflexbending': False, 'nflexedge': False, 'nflexelem': False, 'nflexelemdata': False, 'nflexelemedge': False, 'nflexshelldata': False, + 'nflexstiffness': False, 'nflexvert': False, 'ngeom': False, 'ngravcomp': False,