Import google-deepmind/mujoco_warp from GitHub.

PiperOrigin-RevId: 911361969
Change-Id: I2dd1223a4279b3df2f11adc4f44778b8d153e599
This commit is contained in:
Taylor Howell
2026-05-06 08:47:18 -07:00
committed by Copybara-Service
parent 6376e67070
commit 9ffd50ce6a
9 changed files with 122 additions and 28 deletions
+6 -9
View File
@@ -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)]
+41
View File
@@ -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:
+15 -7
View File
@@ -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,
+16 -1
View File
@@ -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,
+12 -4
View File
@@ -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)
+1
View File
@@ -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"
+17 -4
View File
@@ -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,
-1
View File
@@ -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
+14 -2
View File
@@ -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,