update mjx model and shims

This commit is contained in:
Mustafa Hasekioglu
2026-06-06 14:39:01 -04:00
parent 2e43279ac3
commit 994cfa0e4a
7 changed files with 79 additions and 45 deletions
+8
View File
@@ -755,6 +755,11 @@ class Model(PyTreeNode):
light_pos0: jax.Array
light_dir0: jax.Array
light_cutoff: jax.Array
light_ambient: jax.Array
light_attenuation: jax.Array
light_diffuse: jax.Array
light_exponent: jax.Array
light_specular: jax.Array
mesh_vertadr: np.ndarray
mesh_vertnum: np.ndarray
mesh_faceadr: np.ndarray
@@ -795,6 +800,9 @@ class Model(PyTreeNode):
tex_data: np.ndarray
mat_rgba: jax.Array
mat_texid: jax.Array
mat_emission: jax.Array
mat_specular: jax.Array
mat_shininess: jax.Array
pair_dim: np.ndarray
pair_geom1: np.ndarray
pair_geom2: np.ndarray
+7 -4
View File
@@ -14,17 +14,19 @@
# ==============================================================================
"""DO NOT EDIT. This file is auto-generated."""
import dataclasses
import functools
import jax
import warp as wp
from mujoco.mjx._src import types
import mujoco.mjx.third_party.mujoco_warp as mjwarp
from mujoco.mjx.third_party.mujoco_warp._src import types as mjwp_types
from mujoco.mjx.warp import ffi
from mujoco.mjx.warp.render_context import _MJX_RENDER_CONTEXT_BUFFERS
from mujoco.mjx.warp.render_context import RenderContextPytree
import mujoco.mjx.third_party.mujoco_warp as mjwarp
from mujoco.mjx.third_party.mujoco_warp._src import types as mjwp_types
import warp as wp
_m = mjwarp.Model(
**{f.name: None for f in dataclasses.fields(mjwarp.Model) if f.init}
@@ -48,6 +50,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 _refit_bvh_shim(
# Model
+6 -4
View File
@@ -14,15 +14,17 @@
# ==============================================================================
"""DO NOT EDIT. This file is auto-generated."""
import dataclasses
import functools
import jax
from mujoco.mjx._src import types
from mujoco.mjx.warp import ffi
import mujoco.mjx.third_party.mujoco_warp as mjwarp
from mujoco.mjx.third_party.mujoco_warp._src import types as mjwp_types
import warp as wp
from mujoco.mjx._src import types
import mujoco.mjx.third_party.mujoco_warp as mjwarp
from mujoco.mjx.third_party.mujoco_warp._src import types as mjwp_types
from mujoco.mjx.warp import ffi
_m = mjwarp.Model(
**{f.name: None for f in dataclasses.fields(mjwarp.Model) if f.init}
+6 -4
View File
@@ -14,15 +14,17 @@
# ==============================================================================
"""DO NOT EDIT. This file is auto-generated."""
import dataclasses
import functools
import jax
from mujoco.mjx._src import types
from mujoco.mjx.warp import ffi
import mujoco.mjx.third_party.mujoco_warp as mjwarp
from mujoco.mjx.third_party.mujoco_warp._src import types as mjwp_types
import warp as wp
from mujoco.mjx._src import types
import mujoco.mjx.third_party.mujoco_warp as mjwarp
from mujoco.mjx.third_party.mujoco_warp._src import types as mjwp_types
from mujoco.mjx.warp import ffi
_m = mjwarp.Model(
**{f.name: None for f in dataclasses.fields(mjwarp.Model) if f.init}
+22 -12
View File
@@ -14,17 +14,19 @@
# ==============================================================================
"""DO NOT EDIT. This file is auto-generated."""
import dataclasses
import functools
import jax
import warp as wp
from mujoco.mjx._src import types
import mujoco.mjx.third_party.mujoco_warp as mjwarp
from mujoco.mjx.third_party.mujoco_warp._src import types as mjwp_types
from mujoco.mjx.warp import ffi
from mujoco.mjx.warp.render_context import _MJX_RENDER_CONTEXT_BUFFERS
from mujoco.mjx.warp.render_context import RenderContextPytree
import mujoco.mjx.third_party.mujoco_warp as mjwarp
from mujoco.mjx.third_party.mujoco_warp._src import types as mjwp_types
import warp as wp
_m = mjwarp.Model(
**{f.name: None for f in dataclasses.fields(mjwarp.Model) if f.init}
@@ -171,10 +173,18 @@ def _render_jax_impl(m: types.Model, d: types.Data, ctx: RenderContextPytree):
'geom_xmat',
'geom_xpos',
'light_active',
'light_ambient',
'light_attenuation',
'light_castshadow',
'light_cutoff',
'light_diffuse',
'light_exponent',
'light_specular',
'light_type',
'mat_emission',
'mat_rgba',
'mat_shininess',
'mat_specular',
'mat_texid',
]),
stage_out_argnames=set([]),
@@ -197,18 +207,18 @@ def _render_jax_impl(m: types.Model, d: types.Data, ctx: RenderContextPytree):
m.geom_size,
m.geom_type,
m.light_active,
m._impl.light_ambient,
m._impl.light_attenuation,
m.light_ambient,
m.light_attenuation,
m.light_castshadow,
m.light_cutoff,
m._impl.light_diffuse,
m._impl.light_exponent,
m._impl.light_specular,
m.light_diffuse,
m.light_exponent,
m.light_specular,
m.light_type,
m._impl.mat_emission,
m.mat_emission,
m.mat_rgba,
m._impl.mat_shininess,
m._impl.mat_specular,
m.mat_shininess,
m.mat_specular,
m.mat_texid,
m._impl.mat_texrepeat,
m.mesh_faceadr,
+7 -4
View File
@@ -14,15 +14,17 @@
# ==============================================================================
"""DO NOT EDIT. This file is auto-generated."""
import dataclasses
import functools
import jax
from mujoco.mjx._src import types
from mujoco.mjx.warp import ffi
import mujoco.mjx.third_party.mujoco_warp as mjwarp
from mujoco.mjx.third_party.mujoco_warp._src import types as mjwp_types
import warp as wp
from mujoco.mjx._src import types
import mujoco.mjx.third_party.mujoco_warp as mjwarp
from mujoco.mjx.third_party.mujoco_warp._src import types as mjwp_types
from mujoco.mjx.warp import ffi
_m = mjwarp.Model(
**{f.name: None for f in dataclasses.fields(mjwarp.Model) if f.init}
@@ -46,6 +48,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 _kinematics_shim(
# Model
+23 -17
View File
@@ -15,14 +15,18 @@
"""MJX Warp types.
DO NOT EDIT. This file is auto-generated.
"""
import dataclasses
import typing
from typing import Tuple
import jax
from jax import tree_util
from jax.interpreters import batching
from mujoco.mjx._src import dataclasses as mjx_dataclasses
import numpy as np
from mujoco.mjx._src import dataclasses as mjx_dataclasses
if typing.TYPE_CHECKING:
GraphMode = int
@@ -33,6 +37,7 @@ if typing.TYPE_CHECKING:
else:
try:
from warp._src.jax_experimental.ffi import GraphMode
from mujoco.mjx.third_party.mujoco_warp._src import types as mjwp_types
Callback = mjwp_types.Callback
@@ -41,6 +46,7 @@ else:
Callback = None
PyTreeNode = mjx_dataclasses.PyTreeNode
@dataclasses.dataclass(frozen=True)
@tree_util.register_pytree_node_class
class TileSet:
@@ -52,6 +58,7 @@ class TileSet:
adr: address of each tile in the set
size: size of all the tiles in this set
"""
adr: np.ndarray
size: int
@@ -81,21 +88,18 @@ class BlockDim:
energy_vel_kinetic: energy velocity kinetic block dimension (sensor)
cholesky_factorize: Cholesky factorize block dimension (smooth)
cholesky_solve: Cholesky solve block dimension (smooth)
cholesky_factorize_solve: Cholesky factorize and solve block dimension
(smooth)
cholesky_factorize_solve: Cholesky factorize and solve block dimension (smooth)
solve_LD_sparse_fused: solve LD sparse fused block dimension (smooth)
update_gradient_cholesky: update gradient Cholesky block dimension (solver)
update_gradient_cholesky_blocked: update gradient Cholesky blocked block
dimension (solver)
update_gradient_JTDAJ_sparse: update gradient JTDAJ sparse block dimension
(solver)
update_gradient_JTDAJ_dense: update gradient JTDAJ dense block dimension
(solver)
update_gradient_cholesky_blocked: update gradient Cholesky blocked block dimension (solver)
update_gradient_JTDAJ_sparse: update gradient JTDAJ sparse block dimension (solver)
update_gradient_JTDAJ_dense: update gradient JTDAJ dense block dimension (solver)
linesearch_iterative: linesearch iterative block dimension (solver)
contact_jac_tiled: contact Jacobian tiled block dimension (solver)
qderiv_actuator_dense: qderiv actuator dense block dimension (derivative)
render: render block dimension (render)
"""
actuator_velocity: int
cholesky_factorize: int
cholesky_factorize_solve: int
@@ -127,10 +131,13 @@ class BlockDim:
class StatisticWarp(PyTreeNode):
"""Derived fields from Statistic."""
meaninertia: jax.Array
class OptionWarp(PyTreeNode):
"""Derived fields from Option."""
broadphase: int
broadphase_filter: int
ccd_iterations: int
@@ -146,8 +153,10 @@ class OptionWarp(PyTreeNode):
sdf_iterations: int
sleep_tolerance: jax.Array
class ModelWarp(PyTreeNode):
"""Derived fields from Model."""
D_colind: np.ndarray
D_diag: np.ndarray
D_rowadr: np.ndarray
@@ -232,19 +241,11 @@ class ModelWarp(PyTreeNode):
is_sparse: bool
jnt_limited_ball_adr: np.ndarray
jnt_limited_slide_hinge_adr: np.ndarray
light_ambient: jax.Array
light_attenuation: jax.Array
light_bodyid: np.ndarray
light_diffuse: jax.Array
light_exponent: jax.Array
light_specular: jax.Array
light_targetbodyid: np.ndarray
mapD2M: np.ndarray
mapM2D: np.ndarray
mapM2M: np.ndarray
mat_emission: jax.Array
mat_shininess: jax.Array
mat_specular: jax.Array
mat_texrepeat: jax.Array
max_ten_J_rownnz: int
mesh_polyadr: np.ndarray
@@ -337,8 +338,10 @@ class ModelWarp(PyTreeNode):
wrap_site_adr: np.ndarray
wrap_site_pair_adr: np.ndarray
class DataWarp(PyTreeNode):
"""Derived fields from Data."""
M: jax.Array
actuator_moment: jax.Array
actuator_velocity: jax.Array
@@ -457,6 +460,8 @@ class DataWarp(PyTreeNode):
wrap_obj: jax.Array
wrap_xpos: jax.Array
shape = property(lambda self: self.cacc.shape)
DATA_NON_VMAP = {
'contact__dim',
'contact__dist',
@@ -485,6 +490,7 @@ DATA_NON_VMAP = {
'nworld',
}
def _to_elt(cont, _, d, axis):
return DataWarp(**{
f.name: (