Import Warp v1.14.0
PiperOrigin-RevId: 938464395 Change-Id: Iac99c53566e2e62f6e75f5e0904100919043b463
This commit is contained in:
committed by
Copybara-Service
parent
07e292bc6b
commit
56a182e67c
@@ -20,8 +20,8 @@ jax-cuda12-pjrt==0.8.3; python_version >= '3.13' \
|
||||
jax-cuda12-pjrt==0.5.3; python_version < '3.13' \
|
||||
--hash=sha256:04ee111eaf5fc2692978ad4a5c84d5925e42eb05c1701849ba3a53f6515400cc \
|
||||
--hash=sha256:c5378306568ba0c81b230a779dd3194c9dd10339ab6360ae80928108d37e7f75
|
||||
warp-lang==1.13.0 \
|
||||
--hash=sha256:4375f572301991fe0fbf0af29fc84d76cd27d531432d6df8452b25088c21ea5a \
|
||||
--hash=sha256:ac2479c70ad410d58deb088c2a64168792be588efc1bec22dc51f92238e8c8c3 \
|
||||
--hash=sha256:476b54f0dcf6767f23305a328660a9a74d01e371b6b7c3d77e03e18b4a1bb1a5 \
|
||||
--hash=sha256:47975ea07252d45a4d09d2d1a6cffc55002fa7fbde771c8dceb1baf4b32d4fb8
|
||||
warp-lang==1.14.0 \
|
||||
--hash=sha256:12656050545cc77bf9b9b155399496c1a6279b5b6c59e407507d6858a2beb4a2 \
|
||||
--hash=sha256:70cd127d0e9109417099649fedf9d00f39f1307ccb7a6e9fb87661337868d7de \
|
||||
--hash=sha256:f482787e8da9c9ef045601fde99095e16d604fbcc3cbb4a1e0cef0769388b316 \
|
||||
--hash=sha256:936b49ec78237f9760e58cbe9c46ee6f4244aefbd62071c4fa9fd3b313dfa878
|
||||
|
||||
@@ -531,7 +531,7 @@ def put_model(
|
||||
device: which device to use - if unspecified picks the default device
|
||||
impl: implementation to use
|
||||
graph_mode: CUDA graph capture mode (for Warp only). Use GraphMode enum from
|
||||
warp._src.jax_experimental.ffi. GraphMode.WARP is the default mode.
|
||||
warp._src.jax.ffi.GraphMode.WARP is the default mode.
|
||||
keepalive_refs: optional dict to store references to underlying MuJoCo
|
||||
objects, preventing them from being garbage collected. Required for CPP
|
||||
impl to keep the model alive.
|
||||
|
||||
@@ -29,7 +29,8 @@ from mujoco.mjx.codegen import file
|
||||
import mujoco.mjx.third_party.mujoco_warp as mjwarp
|
||||
import numpy as np
|
||||
import warp as wp
|
||||
from mujoco.mjx.third_party.warp._src.jax_experimental import ffi
|
||||
from warp._src.jax.ffi import FfiArg
|
||||
from warp import JaxCallableGraphMode as GraphMode
|
||||
|
||||
|
||||
_MJX_WARP_TYPES_OUT_FPATH = flags.DEFINE_string(
|
||||
@@ -116,7 +117,7 @@ def _get_target_annotation_node(
|
||||
if annotation in (int, float, bool):
|
||||
return _ast_parse_type(annotation.__name__)
|
||||
|
||||
if annotation is ffi.GraphMode:
|
||||
if annotation is GraphMode:
|
||||
return _ast_parse_type('GraphMode')
|
||||
|
||||
if isinstance(annotation, type) and issubclass(annotation, enum.Enum):
|
||||
@@ -273,7 +274,8 @@ if typing.TYPE_CHECKING:
|
||||
pass
|
||||
else:
|
||||
try:
|
||||
from warp._src.jax_experimental.ffi import GraphMode
|
||||
from mujoco.mjx.third_party.warp._src.jax import ffi as warp_ffi
|
||||
GraphMode = warp_ffi.JaxCallableGraphMode
|
||||
from mujoco.mjx.third_party.mujoco_warp._src import types as mjwp_types
|
||||
Callback = mjwp_types.Callback
|
||||
except ImportError:
|
||||
@@ -501,7 +503,7 @@ def _to_jax_ndim(name: str, wp_type: Any) -> int:
|
||||
if typing.get_args(wp_type)[1] != ...:
|
||||
raise NotImplementedError('Only variadic tuples are supported.')
|
||||
return -1 # signals that dim should be untouched in downstream code.
|
||||
ffi_arg = ffi.FfiArg(name, wp_type)
|
||||
ffi_arg = FfiArg(name, wp_type)
|
||||
return ffi_arg.jax_ndim
|
||||
|
||||
|
||||
@@ -573,7 +575,7 @@ def main(argv):
|
||||
write_core_cls('Statistic', target_fpath, mjx_types_fpath, set_diff=False)
|
||||
write_core_cls(
|
||||
'Option', target_fpath, mjx_types_fpath,
|
||||
extra_annotations={'graph_mode': ffi.GraphMode},
|
||||
extra_annotations={'graph_mode': GraphMode},
|
||||
)
|
||||
write_core_cls('Model', target_fpath, mjx_types_fpath)
|
||||
write_core_cls('Data', target_fpath, mjx_types_fpath, flatten_fields=True)
|
||||
|
||||
+47
-27
@@ -555,6 +555,26 @@ def render(m: Model, d: Data, rc: RenderContext):
|
||||
cast_ray_first_hit = _make_cast_ray(geom_ray_types, first_hit=True)
|
||||
compute_lighting = _make_compute_lighting(cast_ray_first_hit)
|
||||
|
||||
# Extract static parameters to avoid capturing rc and m in the kernel closure.
|
||||
static_use_precomputed_rays = rc.use_precomputed_rays
|
||||
static_znear = rc.znear
|
||||
static_enable_backface_culling = rc.enable_backface_culling
|
||||
static_render_skybox = rc.render_skybox
|
||||
static_skybox_tex_id = rc.skybox_tex_id
|
||||
static_skybox_face_width = rc.skybox_face_width
|
||||
static_use_textures = rc.use_textures
|
||||
static_enable_specular = rc.enable_specular
|
||||
static_enable_emission = rc.enable_emission
|
||||
static_use_ambient_lighting = rc.use_ambient_lighting
|
||||
static_headlight_active = rc.headlight_active
|
||||
static_headlight_ambient = rc.headlight_ambient
|
||||
static_headlight_diffuse = rc.headlight_diffuse
|
||||
static_headlight_specular = rc.headlight_specular
|
||||
static_enable_per_light_ambient = rc.enable_per_light_ambient
|
||||
static_light_attenuation_is_default = rc.light_attenuation_is_default
|
||||
static_has_spot_lights = rc.has_spot_lights
|
||||
static_nlight = m.nlight
|
||||
|
||||
@wp.kernel(module="unique", enable_backward=False)
|
||||
def _render_megakernel(
|
||||
# Model:
|
||||
@@ -649,7 +669,7 @@ def render(m: Model, d: Data, rc: RenderContext):
|
||||
# Map active camera index to MuJoCo camera ID
|
||||
mujoco_cam_id = cam_id_map[camid]
|
||||
|
||||
if wp.static(rc.use_precomputed_rays):
|
||||
if wp.static(static_use_precomputed_rays):
|
||||
ray_dir_local_cam = ray[rayid]
|
||||
else:
|
||||
img_w = cam_res[camid][0]
|
||||
@@ -665,7 +685,7 @@ def render(m: Model, d: Data, rc: RenderContext):
|
||||
img_h,
|
||||
px,
|
||||
py,
|
||||
wp.static(rc.znear),
|
||||
wp.static(static_znear),
|
||||
)
|
||||
|
||||
ray_origin_world = cam_xpos_in[worldid, mujoco_cam_id]
|
||||
@@ -697,7 +717,7 @@ def render(m: Model, d: Data, rc: RenderContext):
|
||||
ray_origin_world,
|
||||
ray_dir_world,
|
||||
float(MJ_MAXVAL),
|
||||
wp.static(rc.enable_backface_culling),
|
||||
wp.static(static_enable_backface_culling),
|
||||
)
|
||||
|
||||
if render_seg[camid] and geom_id != -1:
|
||||
@@ -708,10 +728,10 @@ def render(m: Model, d: Data, rc: RenderContext):
|
||||
|
||||
# Early Out
|
||||
if geom_id == -1:
|
||||
if wp.static(rc.render_skybox) and render_rgb[camid]:
|
||||
if wp.static(static_render_skybox) and render_rgb[camid]:
|
||||
skybox_color = sample_skybox(
|
||||
textures[wp.static(rc.skybox_tex_id)],
|
||||
wp.static(1.0 / float(rc.skybox_face_width)),
|
||||
textures[wp.static(static_skybox_tex_id)],
|
||||
wp.static(1.0 / float(static_skybox_face_width)),
|
||||
ray_dir_world,
|
||||
)
|
||||
rgb_out[worldid, rgb_adr[camid] + rayid_local] = pack_rgba_to_uint32(
|
||||
@@ -745,7 +765,7 @@ def render(m: Model, d: Data, rc: RenderContext):
|
||||
|
||||
base_color = wp.vec3(color[0], color[1], color[2])
|
||||
|
||||
if wp.static(rc.use_textures):
|
||||
if wp.static(static_use_textures):
|
||||
if geom_id != -2:
|
||||
mat_id = geom_matid[worldid % geom_matid.shape[0], geom_id]
|
||||
if mat_id >= 0:
|
||||
@@ -773,27 +793,27 @@ def render(m: Model, d: Data, rc: RenderContext):
|
||||
mat_spec = DEFAULT_MAT_SPECULAR
|
||||
mat_shin_exp = DEFAULT_MAT_SHININESS_EXPONENT
|
||||
mat_emis = DEFAULT_MAT_EMISSION
|
||||
if wp.static(rc.enable_specular or rc.enable_emission):
|
||||
if wp.static(static_enable_specular or static_enable_emission):
|
||||
if geom_id != -2:
|
||||
mat_id_for_spec = geom_matid[worldid % geom_matid.shape[0], geom_id]
|
||||
if mat_id_for_spec >= 0:
|
||||
if wp.static(rc.enable_specular):
|
||||
if wp.static(static_enable_specular):
|
||||
mat_spec = mat_specular[worldid % mat_specular.shape[0], mat_id_for_spec]
|
||||
mat_shin_exp = mat_shininess[worldid % mat_shininess.shape[0], mat_id_for_spec] * MAX_SHININESS
|
||||
if wp.static(rc.enable_emission):
|
||||
if wp.static(static_enable_emission):
|
||||
mat_emis = mat_emission[worldid % mat_emission.shape[0], mat_id_for_spec]
|
||||
|
||||
result = wp.vec3(0.0)
|
||||
if wp.static(rc.enable_emission):
|
||||
if wp.static(static_enable_emission):
|
||||
result = base_color * mat_emis
|
||||
|
||||
if wp.static(rc.use_ambient_lighting):
|
||||
if wp.static(rc.headlight_active):
|
||||
result = result + wp.cw_mul(base_color, wp.static(rc.headlight_ambient))
|
||||
elif wp.static(m.nlight == 0):
|
||||
if wp.static(static_use_ambient_lighting):
|
||||
if wp.static(static_headlight_active):
|
||||
result = result + wp.cw_mul(base_color, wp.static(static_headlight_ambient))
|
||||
elif wp.static(static_nlight == 0):
|
||||
result = result + base_color * NO_LIGHT_AMBIENT_FALLBACK
|
||||
if wp.static(rc.enable_per_light_ambient):
|
||||
for light_index in range(wp.static(m.nlight)):
|
||||
if wp.static(static_enable_per_light_ambient):
|
||||
for light_index in range(wp.static(static_nlight)):
|
||||
if light_active[worldid % light_active.shape[0], light_index]:
|
||||
result = result + wp.cw_mul(base_color, light_ambient[worldid % light_ambient.shape[0], light_index])
|
||||
|
||||
@@ -810,7 +830,7 @@ def render(m: Model, d: Data, rc: RenderContext):
|
||||
light_diffuse_worldid = light_diffuse[worldid % light_diffuse.shape[0]]
|
||||
light_specular_worldid = light_specular[worldid % light_specular.shape[0]]
|
||||
# Apply Lighting for each light
|
||||
for light_index in range(wp.static(m.nlight)):
|
||||
for light_index in range(wp.static(static_nlight)):
|
||||
diff_rgb, spec_rgb = compute_lighting(
|
||||
geom_type,
|
||||
geom_dataid,
|
||||
@@ -849,15 +869,15 @@ def render(m: Model, d: Data, rc: RenderContext):
|
||||
view_dir,
|
||||
mat_spec,
|
||||
mat_shin_exp,
|
||||
wp.static(rc.enable_backface_culling),
|
||||
wp.static(rc.enable_specular),
|
||||
wp.static(rc.light_attenuation_is_default),
|
||||
wp.static(rc.has_spot_lights),
|
||||
wp.static(static_enable_backface_culling),
|
||||
wp.static(static_enable_specular),
|
||||
wp.static(static_light_attenuation_is_default),
|
||||
wp.static(static_has_spot_lights),
|
||||
)
|
||||
result = result + wp.cw_mul(base_color, diff_rgb) + spec_rgb
|
||||
|
||||
# Apply Headlight
|
||||
if wp.static(rc.headlight_active):
|
||||
if wp.static(static_headlight_active):
|
||||
cam_pos = ray_origin_world
|
||||
cam_fwd = -cam_mat_world[:, 2]
|
||||
hl_diff, hl_spec = compute_lighting(
|
||||
@@ -891,15 +911,15 @@ def render(m: Model, d: Data, rc: RenderContext):
|
||||
wp.vec3(1.0, 0.0, 0.0),
|
||||
0.0,
|
||||
0.0,
|
||||
wp.static(rc.headlight_diffuse),
|
||||
wp.static(rc.headlight_specular),
|
||||
wp.static(static_headlight_diffuse),
|
||||
wp.static(static_headlight_specular),
|
||||
normal,
|
||||
hit_point,
|
||||
view_dir,
|
||||
mat_spec,
|
||||
mat_shin_exp,
|
||||
wp.static(rc.enable_backface_culling),
|
||||
wp.static(rc.enable_specular),
|
||||
wp.static(static_enable_backface_culling),
|
||||
wp.static(static_enable_specular),
|
||||
True,
|
||||
False,
|
||||
)
|
||||
|
||||
@@ -0,0 +1,179 @@
|
||||
# SPDX-FileCopyrightText: Copyright (c) 2023 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
import warp
|
||||
|
||||
_wp_module_name_ = "warp.jax"
|
||||
|
||||
|
||||
def device_to_jax(warp_device: warp.DeviceLike):
|
||||
"""Return the Jax device corresponding to a Warp device.
|
||||
|
||||
Returns:
|
||||
:class:`jax.Device`
|
||||
|
||||
Raises:
|
||||
RuntimeError: Failed to find the corresponding Jax device.
|
||||
"""
|
||||
import jax # noqa: PLC0415
|
||||
|
||||
d = warp.get_device(warp_device)
|
||||
|
||||
if d.is_cuda:
|
||||
cuda_devices = jax.devices("cuda")
|
||||
if d.ordinal >= len(cuda_devices):
|
||||
raise RuntimeError(f"Jax device corresponding to '{warp_device}' is not available")
|
||||
return cuda_devices[d.ordinal]
|
||||
else:
|
||||
cpu_devices = jax.devices("cpu")
|
||||
if not cpu_devices:
|
||||
raise RuntimeError(f"Jax device corresponding to '{warp_device}' is not available")
|
||||
return cpu_devices[0]
|
||||
|
||||
|
||||
def device_from_jax(jax_device) -> warp._src.context.Device:
|
||||
"""Return the Warp device corresponding to a Jax device.
|
||||
|
||||
Args:
|
||||
jax_device (jax.Device): A Jax device descriptor.
|
||||
|
||||
Raises:
|
||||
RuntimeError: The Jax device is neither a CPU nor GPU device.
|
||||
"""
|
||||
if jax_device.platform == "cpu":
|
||||
return warp.get_device("cpu")
|
||||
elif jax_device.platform == "gpu":
|
||||
return warp.get_cuda_device(jax_device.id)
|
||||
else:
|
||||
raise RuntimeError(f"Unsupported Jax device platform '{jax_device.platform}'")
|
||||
|
||||
|
||||
def get_jax_device():
|
||||
"""Get the current Jax device."""
|
||||
import jax # noqa: PLC0415
|
||||
|
||||
# TODO: is there a simpler way of getting the Jax "current" device?
|
||||
# check if jax.default_device() context manager is active
|
||||
device = jax.config.jax_default_device
|
||||
# if default device is not set, use first device
|
||||
if device is None:
|
||||
device = jax.local_devices()[0]
|
||||
return device
|
||||
|
||||
|
||||
def dtype_to_jax(warp_dtype):
|
||||
"""Return the Jax dtype corresponding to a Warp dtype.
|
||||
|
||||
Args:
|
||||
warp_dtype: A Warp data type that has a corresponding Jax data type.
|
||||
|
||||
Raises:
|
||||
TypeError: Unable to find a corresponding Jax data type.
|
||||
"""
|
||||
# initialize lookup table on first call to defer jax import
|
||||
if dtype_to_jax.type_map is None:
|
||||
import jax.numpy as jp # noqa: PLC0415
|
||||
|
||||
dtype_to_jax.type_map = {
|
||||
warp.float16: jp.float16,
|
||||
warp.bfloat16: jp.bfloat16,
|
||||
warp.float32: jp.float32,
|
||||
warp.float64: jp.float64,
|
||||
warp.int8: jp.int8,
|
||||
warp.int16: jp.int16,
|
||||
warp.int32: jp.int32,
|
||||
warp.int64: jp.int64,
|
||||
warp.uint8: jp.uint8,
|
||||
warp.uint16: jp.uint16,
|
||||
warp.uint32: jp.uint32,
|
||||
warp.uint64: jp.uint64,
|
||||
warp.bool: jp.bool_,
|
||||
}
|
||||
|
||||
jax_dtype = dtype_to_jax.type_map.get(warp_dtype)
|
||||
if jax_dtype is not None:
|
||||
return jax_dtype
|
||||
else:
|
||||
raise TypeError(f"Cannot convert {warp_dtype} to a Jax type")
|
||||
|
||||
|
||||
def dtype_from_jax(jax_dtype):
|
||||
"""Return the Warp dtype corresponding to a Jax dtype.
|
||||
|
||||
Raises:
|
||||
TypeError: Unable to find a corresponding Warp data type.
|
||||
"""
|
||||
# initialize lookup table on first call to defer jax import
|
||||
if dtype_from_jax.type_map is None:
|
||||
import jax.numpy as jp # noqa: PLC0415
|
||||
|
||||
dtype_from_jax.type_map = {
|
||||
# Jax scalar types
|
||||
jp.float16: warp.float16,
|
||||
jp.bfloat16: warp.bfloat16,
|
||||
jp.float32: warp.float32,
|
||||
jp.float64: warp.float64,
|
||||
jp.int8: warp.int8,
|
||||
jp.int16: warp.int16,
|
||||
jp.int32: warp.int32,
|
||||
jp.int64: warp.int64,
|
||||
jp.uint8: warp.uint8,
|
||||
jp.uint16: warp.uint16,
|
||||
jp.uint32: warp.uint32,
|
||||
jp.uint64: warp.uint64,
|
||||
jp.bool_: warp.bool,
|
||||
# Jax dtype objects
|
||||
jp.dtype(jp.float16): warp.float16,
|
||||
jp.dtype(jp.bfloat16): warp.bfloat16,
|
||||
jp.dtype(jp.float32): warp.float32,
|
||||
jp.dtype(jp.float64): warp.float64,
|
||||
jp.dtype(jp.int8): warp.int8,
|
||||
jp.dtype(jp.int16): warp.int16,
|
||||
jp.dtype(jp.int32): warp.int32,
|
||||
jp.dtype(jp.int64): warp.int64,
|
||||
jp.dtype(jp.uint8): warp.uint8,
|
||||
jp.dtype(jp.uint16): warp.uint16,
|
||||
jp.dtype(jp.uint32): warp.uint32,
|
||||
jp.dtype(jp.uint64): warp.uint64,
|
||||
jp.dtype(jp.bool_): warp.bool,
|
||||
}
|
||||
|
||||
wp_dtype = dtype_from_jax.type_map.get(jax_dtype)
|
||||
if wp_dtype is not None:
|
||||
return wp_dtype
|
||||
else:
|
||||
raise TypeError(f"Cannot convert {jax_dtype} to a Warp type")
|
||||
|
||||
|
||||
# lookup tables initialized when needed
|
||||
dtype_from_jax.type_map = None
|
||||
dtype_to_jax.type_map = None
|
||||
|
||||
|
||||
def to_jax(warp_array):
|
||||
"""
|
||||
Convert a Warp array to a Jax array without copying the data.
|
||||
|
||||
Args:
|
||||
warp_array (warp.array): The Warp array to convert.
|
||||
|
||||
Returns:
|
||||
jax.Array: The converted Jax array.
|
||||
"""
|
||||
import jax.dlpack # noqa: PLC0415
|
||||
|
||||
return jax.dlpack.from_dlpack(warp_array)
|
||||
|
||||
|
||||
def from_jax(jax_array, dtype=None) -> warp.array:
|
||||
"""Convert a Jax array to a Warp array without copying the data.
|
||||
|
||||
Args:
|
||||
jax_array (jax.Array): The Jax array to convert.
|
||||
dtype: The target data type of the resulting Warp array. Defaults to the Jax array's data type mapped to a Warp data type.
|
||||
|
||||
Returns:
|
||||
warp.array: The converted Warp array.
|
||||
"""
|
||||
|
||||
return warp.from_dlpack(jax_array, dtype=dtype)
|
||||
+11
-11
@@ -5,12 +5,12 @@ import ctypes
|
||||
from functools import reduce
|
||||
|
||||
import warp as wp
|
||||
from warp._src.context import type_str
|
||||
from warp._src.jax import get_jax_device
|
||||
from warp._src.types import array_t, launch_bounds_t, matches_array_class, strides_from_shape
|
||||
from warp._src.utils import warn
|
||||
from warp._src.context import _build_launch_bounds, type_str
|
||||
from mujoco.mjx.third_party.warp._src.jax import get_jax_device
|
||||
from warp._src.logger import log_warning
|
||||
from warp._src.types import array_t, matches_array_class, strides_from_shape
|
||||
|
||||
_wp_module_name_ = "warp.jax_experimental.custom_call"
|
||||
_wp_module_name_ = "warp.jax.custom_call"
|
||||
|
||||
_jax_warp_p = None
|
||||
|
||||
@@ -25,11 +25,11 @@ def jax_kernel(kernel, launch_dims=None, quiet=False):
|
||||
|
||||
.. deprecated:: 1.10.0
|
||||
This version of ``jax_kernel()`` is deprecated for JAX >= 0.5.0 and is not supported
|
||||
with JAX >= 0.8.0. Use :func:`warp.jax_experimental.ffi.jax_kernel` instead, which
|
||||
with JAX >= 0.8.0. Use :func:`warp.jax_kernel` instead, which
|
||||
is the default implementation as of Warp 1.10.
|
||||
|
||||
This implementation requires JAX version 0.4.25 - 0.7.x. For JAX 0.8.0 and later,
|
||||
use the FFI-based implementation at :func:`warp.jax_experimental.ffi.jax_kernel`.
|
||||
use the FFI-based implementation at :func:`warp.jax_kernel`.
|
||||
|
||||
Args:
|
||||
kernel: The Warp kernel to be wrapped.
|
||||
@@ -56,14 +56,14 @@ def jax_kernel(kernel, launch_dims=None, quiet=False):
|
||||
f"but installed JAX version is {jax.__version_info__}."
|
||||
)
|
||||
if jax.__version_info__ >= (0, 8, 0):
|
||||
msg += " Please use warp.jax_experimental.ffi.jax_kernel instead."
|
||||
msg += " Please use warp.jax_kernel instead."
|
||||
raise RuntimeError(msg)
|
||||
|
||||
# deprecation warning
|
||||
if jax.__version_info__ >= (0, 5, 0) and not quiet:
|
||||
warn(
|
||||
log_warning(
|
||||
"This version of jax_kernel() is deprecated and will not be supported with newer JAX versions. "
|
||||
"Please use the newer FFI version instead (warp.jax_experimental.ffi.jax_kernel). "
|
||||
"Please use the newer FFI version instead (warp.jax_kernel). "
|
||||
"As of Warp release 1.10, the FFI version is the default implementation of jax_kernel(). "
|
||||
"Pass quiet=True to disable this warning.",
|
||||
DeprecationWarning,
|
||||
@@ -97,7 +97,7 @@ def _warp_custom_callback(stream, buffers, opaque, opaque_len):
|
||||
|
||||
# Parse launch dimensions.
|
||||
dims = [int(d) for d in dim_str.split(",")]
|
||||
bounds = launch_bounds_t(dims)
|
||||
bounds = _build_launch_bounds(dims, kernel.adj.kernel_dim)
|
||||
|
||||
# Parse arguments.
|
||||
arg_strings = args_str.split(";")
|
||||
+131
-72
@@ -11,25 +11,22 @@ import traceback
|
||||
from collections.abc import Callable
|
||||
from enum import IntEnum
|
||||
|
||||
import jax
|
||||
|
||||
import warp as wp
|
||||
from warp._src.codegen import get_full_arg_spec, make_full_qualified_name
|
||||
from warp._src.context import CudaMemcpyKind
|
||||
from warp._src.jax import get_jax_device
|
||||
from warp._src.context import CudaMemcpyKind, _build_launch_bounds
|
||||
from mujoco.mjx.third_party.warp._src.jax import get_jax_device
|
||||
from warp._src.logger import log_warning
|
||||
from warp._src.types import (
|
||||
array_t,
|
||||
launch_bounds_t,
|
||||
matches_array_class,
|
||||
strides_from_shape,
|
||||
type_size_in_bytes,
|
||||
type_to_warp,
|
||||
)
|
||||
from warp._src.utils import warn
|
||||
|
||||
from .xla_ffi import *
|
||||
|
||||
_wp_module_name_ = "warp.jax_experimental.ffi"
|
||||
_wp_module_name_ = "warp.jax.ffi"
|
||||
|
||||
# Holders for the custom callbacks to keep them alive.
|
||||
_FFI_KERNEL_REGISTRY: dict[tuple, FfiKernel] = {}
|
||||
@@ -41,8 +38,20 @@ _FFI_REGISTRY_LOCK = threading.Lock()
|
||||
# Lock when XLA invokes callbacks from multiple threads.
|
||||
_FFI_CALLBACK_LOCK = threading.Lock()
|
||||
|
||||
# Sentinel for detecting when per-call kwargs are passed to differentiable wrappers.
|
||||
_MISSING = object()
|
||||
|
||||
JAX_CALLABLE_DEFAULT_GRAPH_CACHE_MAX = 32
|
||||
|
||||
|
||||
def _get_jax():
|
||||
import jax # noqa: PLC0415
|
||||
|
||||
return jax
|
||||
|
||||
|
||||
def check_jax_version():
|
||||
jax = _get_jax()
|
||||
# check if JAX version supports this
|
||||
if jax.__version_info__ < (0, 5, 0):
|
||||
msg = (
|
||||
@@ -69,8 +78,8 @@ def compute_batch_size(shape, batch_ndim):
|
||||
return batch_size
|
||||
|
||||
|
||||
class GraphMode(IntEnum):
|
||||
"""CUDA graph capture modes for :func:`warp.jax_experimental.jax_callable`.
|
||||
class JaxCallableGraphMode(IntEnum):
|
||||
"""CUDA graph capture modes for :func:`warp.jax_callable`.
|
||||
|
||||
These modes control whether JAX or Warp captures a CUDA graph, and whether
|
||||
staging buffers are used when capturing with Warp.
|
||||
@@ -88,12 +97,20 @@ class GraphMode(IntEnum):
|
||||
"""Capture a Warp graph using staging buffers and perform memcpy outside the graph."""
|
||||
|
||||
|
||||
class ModulePreloadMode(IntEnum):
|
||||
GraphMode = JaxCallableGraphMode
|
||||
|
||||
|
||||
class JaxModulePreloadMode(IntEnum):
|
||||
"""Module preload modes for JAX interop callables."""
|
||||
|
||||
NONE = 0 # don't preload modules
|
||||
CURRENT_DEVICE = 1 # preload on currently active device
|
||||
ALL_DEVICES = 2 # preload on all supported devices
|
||||
|
||||
|
||||
ModulePreloadMode = JaxModulePreloadMode
|
||||
|
||||
|
||||
class FfiArg:
|
||||
def __init__(self, name, type, in_out=False):
|
||||
self.name = name
|
||||
@@ -211,6 +228,7 @@ class FfiKernel:
|
||||
self.input_output_aliases = input_output_aliases
|
||||
|
||||
# register the callback
|
||||
jax = _get_jax()
|
||||
FFI_CCALLFUNC = ctypes.CFUNCTYPE(
|
||||
ctypes.c_void_p, ctypes.POINTER(XLA_FFI_CallFrame)
|
||||
)
|
||||
@@ -226,6 +244,7 @@ class FfiKernel:
|
||||
jax.ffi.register_ffi_target(self.name, ffi_capsule_host, platform="Host")
|
||||
|
||||
def __call__(self, *args, output_dims=None, launch_dims=None, vmap_method=None):
|
||||
jax = _get_jax()
|
||||
num_inputs = len(args)
|
||||
if num_inputs != self.num_inputs:
|
||||
raise ValueError(f"Expected {self.num_inputs} inputs, but got {num_inputs}")
|
||||
@@ -314,19 +333,19 @@ class FfiKernel:
|
||||
)
|
||||
|
||||
# preload on the specified devices
|
||||
if self.module_preload_mode == ModulePreloadMode.CURRENT_DEVICE:
|
||||
if self.module_preload_mode == JaxModulePreloadMode.CURRENT_DEVICE:
|
||||
device = wp.device_from_jax(get_jax_device())
|
||||
self.kernel.module.load(device)
|
||||
elif self.module_preload_mode == ModulePreloadMode.ALL_DEVICES:
|
||||
elif self.module_preload_mode == JaxModulePreloadMode.ALL_DEVICES:
|
||||
for d in jax.local_devices():
|
||||
try:
|
||||
dev = wp.device_from_jax(d)
|
||||
# we only support CUDA devices for now
|
||||
if dev.is_cuda:
|
||||
self.kernel.module.load(dev)
|
||||
except Exception:
|
||||
# ignore unsupported devices like TPUs
|
||||
pass
|
||||
# we only support CUDA devices for now
|
||||
if dev.is_cuda:
|
||||
self.kernel.module.load(dev)
|
||||
|
||||
# save launch data to be retrieved by callback
|
||||
launch_id = self.launch_id
|
||||
@@ -429,7 +448,7 @@ class FfiKernel:
|
||||
# roll batch size into the first launch dimension
|
||||
launch_dims = (batch_size * launch_dims[0], *launch_dims[1:])
|
||||
|
||||
launch_bounds = launch_bounds_t(launch_dims)
|
||||
launch_bounds = _build_launch_bounds(launch_dims, self.kernel.adj.kernel_dim)
|
||||
kernel_params[0] = ctypes.addressof(launch_bounds)
|
||||
|
||||
# get device and stream
|
||||
@@ -500,7 +519,7 @@ class FfiCallDesc:
|
||||
|
||||
|
||||
class FfiCallable:
|
||||
default_graph_cache_max: int | None = 32
|
||||
default_graph_cache_max: int | None = JAX_CALLABLE_DEFAULT_GRAPH_CACHE_MAX
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
@@ -616,6 +635,7 @@ class FfiCallable:
|
||||
self.input_output_aliases = input_output_aliases
|
||||
|
||||
# register the callback
|
||||
jax = _get_jax()
|
||||
FFI_CCALLFUNC = ctypes.CFUNCTYPE(
|
||||
ctypes.c_void_p, ctypes.POINTER(XLA_FFI_CallFrame)
|
||||
)
|
||||
@@ -631,6 +651,7 @@ class FfiCallable:
|
||||
jax.ffi.register_ffi_target(self.name, ffi_capsule_host, platform="Host")
|
||||
|
||||
def __call__(self, *args, output_dims=None, vmap_method=None):
|
||||
jax = _get_jax()
|
||||
num_inputs = len(args)
|
||||
if num_inputs != self.num_inputs:
|
||||
input_names = ", ".join(arg.name for arg in self.input_args)
|
||||
@@ -709,18 +730,18 @@ class FfiCallable:
|
||||
# preload on the specified devices
|
||||
# NOTE: if the target function uses kernels from different modules, they will not be loaded here
|
||||
module = wp.get_module(self.func.__module__)
|
||||
if self.module_preload_mode == ModulePreloadMode.CURRENT_DEVICE:
|
||||
if self.module_preload_mode == JaxModulePreloadMode.CURRENT_DEVICE:
|
||||
device = wp.device_from_jax(get_jax_device())
|
||||
module.load(device)
|
||||
elif self.module_preload_mode == ModulePreloadMode.ALL_DEVICES:
|
||||
elif self.module_preload_mode == JaxModulePreloadMode.ALL_DEVICES:
|
||||
for d in jax.local_devices():
|
||||
try:
|
||||
dev = wp.device_from_jax(d)
|
||||
if dev.is_cuda or dev.is_cpu:
|
||||
module.load(dev)
|
||||
except Exception:
|
||||
# ignore unsupported devices like TPUs
|
||||
pass
|
||||
if dev.is_cuda or dev.is_cpu:
|
||||
module.load(dev)
|
||||
|
||||
# save call data to be retrieved by callback
|
||||
call_id = self.call_id
|
||||
@@ -741,7 +762,7 @@ class FfiCallable:
|
||||
metadata_ext.contents.metadata.contents.api_version.major_version = 0
|
||||
metadata_ext.contents.metadata.contents.api_version.minor_version = 1
|
||||
# Turn on CUDA graphs for this handler if on CUDA platform.
|
||||
if self.graph_mode is GraphMode.JAX and platform == "CUDA":
|
||||
if self.graph_mode is JaxCallableGraphMode.JAX and platform == "CUDA":
|
||||
metadata_ext.contents.metadata.contents.traits = (
|
||||
XLA_FFI_Handler_TraitsBits.COMMAND_BUFFER_COMPATIBLE
|
||||
)
|
||||
@@ -800,7 +821,7 @@ class FfiCallable:
|
||||
cuda_stream = get_stream_from_callframe(call_frame.contents)
|
||||
device_ordinal = get_device_ordinal_from_callframe(call_frame.contents)
|
||||
|
||||
if self.graph_mode == GraphMode.WARP:
|
||||
if self.graph_mode == JaxCallableGraphMode.WARP:
|
||||
# check if we already captured an identical call
|
||||
ip = [inputs[i].contents.data for i in self.array_input_indices]
|
||||
op = [outputs[i].contents.data for i in self.array_output_indices]
|
||||
@@ -819,7 +840,7 @@ class FfiCallable:
|
||||
# early out
|
||||
return
|
||||
|
||||
elif self.graph_mode == GraphMode.WARP_STAGED_EX:
|
||||
elif self.graph_mode == JaxCallableGraphMode.WARP_STAGED_EX:
|
||||
if call_desc.capture is not None:
|
||||
graph_exec = call_desc.capture.graph.graph_exec
|
||||
context = call_desc.capture.graph.device.context
|
||||
@@ -866,7 +887,7 @@ class FfiCallable:
|
||||
# early out
|
||||
return
|
||||
|
||||
elif self.graph_mode == GraphMode.WARP_STAGED:
|
||||
elif self.graph_mode == JaxCallableGraphMode.WARP_STAGED:
|
||||
if call_desc.capture is not None:
|
||||
graph_exec = call_desc.capture.graph.graph_exec
|
||||
|
||||
@@ -938,7 +959,7 @@ class FfiCallable:
|
||||
# keep a reference to the capture object to prevent required modules getting unloaded
|
||||
call_desc.capture = capture
|
||||
|
||||
elif self.graph_mode == GraphMode.WARP and device.is_cuda:
|
||||
elif self.graph_mode == JaxCallableGraphMode.WARP and device.is_cuda:
|
||||
# capturing with WARP
|
||||
with wp.ScopedCapture() as capture:
|
||||
self.func(*arg_list)
|
||||
@@ -951,7 +972,7 @@ class FfiCallable:
|
||||
if self._graph_cache_max is not None and len(self.captures) > self._graph_cache_max:
|
||||
self.captures.popitem(last=False)
|
||||
|
||||
elif self.graph_mode == GraphMode.WARP_STAGED_EX and device.is_cuda:
|
||||
elif self.graph_mode == JaxCallableGraphMode.WARP_STAGED_EX and device.is_cuda:
|
||||
# capturing with WARP using staging buffers and memcopies done outside of the graph
|
||||
wp_memcpy_batch = wp._src.context.runtime.core.wp_memcpy_batch
|
||||
|
||||
@@ -994,7 +1015,7 @@ class FfiCallable:
|
||||
# TODO: we should have a way of freeing this
|
||||
call_desc.capture = capture
|
||||
|
||||
elif self.graph_mode == GraphMode.WARP_STAGED and device.is_cuda:
|
||||
elif self.graph_mode == JaxCallableGraphMode.WARP_STAGED and device.is_cuda:
|
||||
# capturing with WARP using staging buffers and memcopies done inside of the graph
|
||||
wp_cuda_graph_insert_memcpy_batch = (
|
||||
wp._src.context.runtime.core.wp_cuda_graph_insert_memcpy_batch
|
||||
@@ -1198,7 +1219,7 @@ def jax_kernel(
|
||||
launch_dims=None,
|
||||
output_dims=None,
|
||||
in_out_argnames=None,
|
||||
module_preload_mode=ModulePreloadMode.CURRENT_DEVICE,
|
||||
module_preload_mode=JaxModulePreloadMode.CURRENT_DEVICE,
|
||||
enable_backward: bool = False,
|
||||
has_side_effect: bool = False,
|
||||
):
|
||||
@@ -1212,12 +1233,16 @@ def jax_kernel(
|
||||
This must include the number of ``in_out_arguments``.
|
||||
vmap_method: String specifying how the callback transforms under ``vmap()``.
|
||||
This argument can also be specified for individual calls.
|
||||
launch_dims: Specify the default kernel launch dimensions. If None, launch
|
||||
dimensions are inferred from the shape of the first array argument.
|
||||
This argument can also be specified for individual calls.
|
||||
launch_dims: Specify the kernel launch dimensions. If None, launch dimensions
|
||||
are inferred from the shape of the first array argument. When
|
||||
``enable_backward=False``, this value will be used by default but
|
||||
can be overridden for individual calls. When ``enable_backward=True``,
|
||||
this value is fixed at construction time and cannot be overridden
|
||||
per call.
|
||||
output_dims: Specify the default dimensions of output arrays. If None, output
|
||||
dimensions are inferred from the launch dimensions.
|
||||
This argument can also be specified for individual calls.
|
||||
Not supported when ``enable_backward=True``.
|
||||
in_out_argnames: Names of arguments that are both inputs and outputs (aliased buffers).
|
||||
These must be array arguments that appear before any pure output arguments in the
|
||||
kernel signature. The number of in-out arguments is included in ``num_outputs``.
|
||||
@@ -1233,9 +1258,11 @@ def jax_kernel(
|
||||
- Input and input-output arguments must precede the output arguments in the ``kernel`` definition.
|
||||
- There must be at least one output or input-output argument.
|
||||
- Only the CUDA backend is supported.
|
||||
- ``output_dims`` and ``in_out_argnames`` are not supported when ``enable_backward=True``.
|
||||
"""
|
||||
|
||||
check_jax_version()
|
||||
jax = _get_jax()
|
||||
|
||||
if isinstance(output_dims, dict):
|
||||
hashable_output_dims = tuple(sorted(output_dims.items()))
|
||||
@@ -1283,11 +1310,11 @@ def jax_kernel(
|
||||
"jax_kernel(): Input-output arguments (in_out_argnames) are not supported when enable_backward=True."
|
||||
)
|
||||
|
||||
# TODO: we should support passing these to the forward and backward callables
|
||||
if launch_dims is not None or output_dims is not None:
|
||||
raise NotImplementedError(
|
||||
"jax_kernel(): Custom dimensions (launch_dims, output_dims) are not supported when enable_backward=True."
|
||||
)
|
||||
# TODO: support output_dims with enable_backward=True (requires separate
|
||||
# output-buffer allocation logic). launch_dims is supported below: the
|
||||
# captured value is applied to both the forward and the adjoint launches.
|
||||
if output_dims is not None:
|
||||
raise NotImplementedError("jax_kernel(): output_dims is not yet supported when enable_backward=True.")
|
||||
|
||||
# Differentiable path: build a custom VJP wrapper inline.
|
||||
# Infer the original kernel signature (names and annotations)
|
||||
@@ -1307,8 +1334,17 @@ def jax_kernel(
|
||||
else:
|
||||
raise TypeError(f"Invalid type for argument '{p.name}', expected array or scalar, got {type}")
|
||||
|
||||
# Capture an explicit user-supplied launch_dims so the same value is used
|
||||
# for both the forward launch and the adjoint launch (required for
|
||||
# correct gradient values when array.ndim > kernel.tid_ndim).
|
||||
# Reuse `hashable_launch_dims` (computed above for the cache key path)
|
||||
# so 1-D integer and sequence forms are normalized identically.
|
||||
_user_launch_dims = hashable_launch_dims if launch_dims is not None else None
|
||||
|
||||
def _resolve_launch_dims(call_args):
|
||||
# determine launch dimensions from the shape of the first input array
|
||||
if _user_launch_dims is not None:
|
||||
return _user_launch_dims
|
||||
# Fallback: determine launch dimensions from the shape of the first input array
|
||||
for i, p in enumerate(parameters[:num_inputs]):
|
||||
param_type = p.annotation
|
||||
if matches_array_class(param_type, wp.array):
|
||||
@@ -1360,11 +1396,15 @@ def jax_kernel(
|
||||
try:
|
||||
gi.zero_()
|
||||
except Exception as e:
|
||||
warn(f"Failed to zero gradient array: {e}", stacklevel=2)
|
||||
log_warning(f"Failed to zero gradient array: {e}", stacklevel=2)
|
||||
raise e
|
||||
|
||||
# NOTE: We cannot use a passed launch_dims here, the backward rule doesn't receive it (and it could be wrong under pmap/vmap).
|
||||
# We need to infer from the inputs.
|
||||
# The same _resolve_launch_dims() is used here so that the adjoint
|
||||
# kernel launches with exactly the same iteration space as the forward
|
||||
# kernel (captured from the enclosing scope via _user_launch_dims when
|
||||
# the caller supplied an explicit value, otherwise inferred from the
|
||||
# inputs). This matches the forward path and avoids N x
|
||||
# over-accumulation in atomic_add when array.ndim > kernel.tid_ndim.
|
||||
wp.launch(
|
||||
kernel,
|
||||
dim=_resolve_launch_dims(inputs),
|
||||
@@ -1498,46 +1538,65 @@ def jax_kernel(
|
||||
vmap_method,
|
||||
module_preload_mode,
|
||||
has_side_effect,
|
||||
# Include the normalized launch_dims so that wrapping the same kernel
|
||||
# with different launch_dims produces independent cache entries.
|
||||
# Reusing _user_launch_dims ensures int and 1-tuple forms of the same
|
||||
# value map to the same key.
|
||||
_user_launch_dims,
|
||||
)
|
||||
|
||||
if static_args:
|
||||
static_names = [parameters[i].name for i in static_args]
|
||||
else:
|
||||
static_names = []
|
||||
|
||||
def _user_callable(*args):
|
||||
return jax_func(*args)
|
||||
key = (*key, tuple(sorted(static_names)))
|
||||
|
||||
_user_callable.__signature__ = signature
|
||||
def _user_callable(*args):
|
||||
return jax_func(*args)
|
||||
|
||||
# Cache differentiable wrapper
|
||||
key = (*key, tuple(sorted(static_names)))
|
||||
with _FFI_REGISTRY_LOCK:
|
||||
cached = _FFI_DIFF_KERNEL_REGISTRY.get(key)
|
||||
if cached is None:
|
||||
cached = jax.jit(_user_callable, static_argnames=tuple(static_names))
|
||||
_FFI_DIFF_KERNEL_REGISTRY[key] = cached
|
||||
return _FFI_DIFF_KERNEL_REGISTRY[key]
|
||||
_user_callable.__signature__ = signature
|
||||
|
||||
# Cache differentiable wrapper (no static args)
|
||||
key = (*key, ())
|
||||
# Cache differentiable wrapper
|
||||
with _FFI_REGISTRY_LOCK:
|
||||
cached = _FFI_DIFF_KERNEL_REGISTRY.get(key)
|
||||
if cached is None:
|
||||
_FFI_DIFF_KERNEL_REGISTRY[key] = jax_func
|
||||
cached = jax_func
|
||||
return cached
|
||||
cached = jax.jit(_user_callable, static_argnames=tuple(static_names))
|
||||
_FFI_DIFF_KERNEL_REGISTRY[key] = cached
|
||||
|
||||
# Thin Python-level wrapper that intercepts FfiKernel-style per-call kwargs
|
||||
# and raises informative errors before JAX tracing begins.
|
||||
def _checked_wrapper(*args, launch_dims=_MISSING, output_dims=_MISSING, vmap_method=_MISSING):
|
||||
if launch_dims is not _MISSING:
|
||||
raise TypeError(
|
||||
"jax_kernel(): launch_dims cannot be overridden per-call when enable_backward=True "
|
||||
f"(this wrapper was created with launch_dims={_user_launch_dims!r}). "
|
||||
"Call jax_kernel() again with a different launch_dims to create a new wrapper."
|
||||
)
|
||||
if output_dims is not _MISSING:
|
||||
raise TypeError("jax_kernel(): output_dims is not supported when enable_backward=True.")
|
||||
if vmap_method is not _MISSING:
|
||||
raise TypeError(
|
||||
"jax_kernel(): vmap_method cannot be overridden per-call when enable_backward=True; "
|
||||
"it is fixed at construction time. "
|
||||
"Call jax_kernel() again with a different vmap_method to create a new wrapper."
|
||||
)
|
||||
return cached(*args)
|
||||
|
||||
return _checked_wrapper
|
||||
|
||||
|
||||
def jax_callable(
|
||||
func: Callable,
|
||||
num_outputs: int = 1,
|
||||
graph_mode: GraphMode = GraphMode.JAX,
|
||||
graph_mode: JaxCallableGraphMode = JaxCallableGraphMode.JAX,
|
||||
vmap_method: str | None = "broadcast_all",
|
||||
output_dims=None,
|
||||
in_out_argnames=None,
|
||||
stage_in_argnames=None,
|
||||
stage_out_argnames=None,
|
||||
graph_cache_max: int | None = None,
|
||||
module_preload_mode: ModulePreloadMode = ModulePreloadMode.CURRENT_DEVICE,
|
||||
graph_cache_max: int | None = JAX_CALLABLE_DEFAULT_GRAPH_CACHE_MAX,
|
||||
module_preload_mode: JaxModulePreloadMode = JaxModulePreloadMode.CURRENT_DEVICE,
|
||||
has_side_effect: bool = False,
|
||||
):
|
||||
"""Create a JAX callback from an annotated Python function.
|
||||
@@ -1551,10 +1610,10 @@ def jax_callable(
|
||||
num_outputs: Specify the number of output arguments if greater than 1.
|
||||
This must include the number of ``in_out_arguments``.
|
||||
graph_mode: CUDA graph capture mode.
|
||||
``GraphMode.JAX`` (default): Let JAX capture the graph, which may be used as a subgraph in an enclosing JAX capture.
|
||||
``GraphMode.WARP``: Let Warp capture the graph. Use this mode when the callable cannot be used as a subgraph,
|
||||
``JaxCallableGraphMode.JAX`` (default): Let JAX capture the graph, which may be used as a subgraph in an enclosing JAX capture.
|
||||
``JaxCallableGraphMode.WARP``: Let Warp capture the graph. Use this mode when the callable cannot be used as a subgraph,
|
||||
such as when the callable uses conditional graph nodes.
|
||||
``GraphMode.NONE``: Disable graph capture. Use when the callable performs operations that are not legal in a graph,
|
||||
``JaxCallableGraphMode.NONE``: Disable graph capture. Use when the callable performs operations that are not legal in a graph,
|
||||
such as host synchronization.
|
||||
vmap_method: String specifying how the callback transforms under ``vmap()``.
|
||||
This argument can also be specified for individual calls.
|
||||
@@ -1564,12 +1623,12 @@ def jax_callable(
|
||||
in_out_argnames: Names of arguments that are both inputs and outputs (aliased buffers).
|
||||
These must be array arguments that appear before any pure output arguments in the
|
||||
function signature. The number of in-out arguments is included in ``num_outputs``.
|
||||
stage_in_argnames: Names of input arguments that need to be copied with ``GraphMode.WARP_STAGED*``.
|
||||
stage_in_argnames: Names of input arguments that need to be copied with ``JaxCallableGraphMode.WARP_STAGED*``.
|
||||
If ``None``, copy all input arguments.
|
||||
stage_out_argnames: Names of output arguments that need to be copied with ``GraphMode.WARP_STAGED*``.
|
||||
stage_out_argnames: Names of output arguments that need to be copied with ``JaxCallableGraphMode.WARP_STAGED*``.
|
||||
If ``None``, copy all output arguments.
|
||||
graph_cache_max: Maximum number of cached graphs captured using ``GraphMode.WARP``.
|
||||
If ``None``, use ``warp.jax_experimental.get_jax_callable_default_graph_cache_max()``.
|
||||
graph_cache_max: Maximum number of cached graphs captured using ``JaxCallableGraphMode.WARP``.
|
||||
If ``None``, the graph cache is unlimited.
|
||||
module_preload_mode: Specify the devices where the module should be preloaded.
|
||||
has_side_effect: Whether the custom call has side effects. When True,
|
||||
the FFI call will be executed even when the outputs are not used.
|
||||
@@ -1584,9 +1643,6 @@ def jax_callable(
|
||||
|
||||
check_jax_version()
|
||||
|
||||
if graph_cache_max is None:
|
||||
graph_cache_max = FfiCallable.default_graph_cache_max
|
||||
|
||||
if isinstance(output_dims, dict):
|
||||
hashable_output_dims = tuple(sorted(output_dims.items()))
|
||||
elif hasattr(output_dims, "__len__"):
|
||||
@@ -1631,14 +1687,14 @@ def jax_callable(
|
||||
|
||||
def get_jax_callable_default_graph_cache_max():
|
||||
"""
|
||||
Get the maximum size of the graph cache for graphs captured using ``GraphMode.WARP``, unlimited if ``None``.
|
||||
Get the maximum size of the graph cache for graphs captured using ``JaxCallableGraphMode.WARP``, unlimited if ``None``.
|
||||
"""
|
||||
return FfiCallable.default_graph_cache_max
|
||||
|
||||
|
||||
def set_jax_callable_default_graph_cache_max(cache_max: int | None):
|
||||
"""
|
||||
Set the maximum size of the graph cache for graphs captured using ``GraphMode.WARP``, unlimited if ``None``.
|
||||
Set the maximum size of the graph cache for graphs captured using ``JaxCallableGraphMode.WARP``, unlimited if ``None``.
|
||||
"""
|
||||
FfiCallable.default_graph_cache_max = cache_max
|
||||
|
||||
@@ -1677,6 +1733,7 @@ def register_ffi_callback(name: str, func: Callable, graph_compatible: bool = Tr
|
||||
"""
|
||||
|
||||
check_jax_version()
|
||||
jax = _get_jax()
|
||||
|
||||
# TODO check that the name is not already registered
|
||||
|
||||
@@ -1765,6 +1822,8 @@ def get_warp_shape(arg, dims):
|
||||
|
||||
|
||||
def get_jax_output_type(arg, dims):
|
||||
jax = _get_jax()
|
||||
|
||||
if isinstance(dims, int):
|
||||
dims = (dims,)
|
||||
|
||||
+71
-45
@@ -3,13 +3,78 @@
|
||||
|
||||
import ctypes
|
||||
import enum
|
||||
import threading
|
||||
|
||||
import jax.numpy as jnp
|
||||
import numpy as np
|
||||
|
||||
import warp as wp
|
||||
|
||||
_wp_module_name_ = "warp.jax_experimental.xla_ffi"
|
||||
_wp_module_name_ = "warp.jax.xla_ffi"
|
||||
|
||||
_xla_data_type_to_constructor = None
|
||||
_XLA_DATA_TYPE_TO_CONSTRUCTOR_LOCK = threading.Lock()
|
||||
|
||||
|
||||
def _get_jnp():
|
||||
import jax.numpy as jnp # noqa: PLC0415
|
||||
|
||||
return jnp
|
||||
|
||||
|
||||
def _get_xla_data_type_to_constructor():
|
||||
global _xla_data_type_to_constructor
|
||||
|
||||
if _xla_data_type_to_constructor is None:
|
||||
with _XLA_DATA_TYPE_TO_CONSTRUCTOR_LOCK:
|
||||
if _xla_data_type_to_constructor is None:
|
||||
jnp = _get_jnp()
|
||||
constructors = {
|
||||
# XLA_FFI_DataType.INVALID
|
||||
XLA_FFI_DataType.PRED: jnp.bool,
|
||||
XLA_FFI_DataType.S8: jnp.int8,
|
||||
XLA_FFI_DataType.S16: jnp.int16,
|
||||
XLA_FFI_DataType.S32: jnp.int32,
|
||||
XLA_FFI_DataType.S64: jnp.int64,
|
||||
XLA_FFI_DataType.U8: jnp.uint8,
|
||||
XLA_FFI_DataType.U16: jnp.uint16,
|
||||
XLA_FFI_DataType.U32: jnp.uint32,
|
||||
XLA_FFI_DataType.U64: jnp.uint64,
|
||||
XLA_FFI_DataType.F16: jnp.float16,
|
||||
XLA_FFI_DataType.F32: jnp.float32,
|
||||
XLA_FFI_DataType.F64: jnp.float64,
|
||||
XLA_FFI_DataType.BF16: jnp.bfloat16,
|
||||
XLA_FFI_DataType.C64: jnp.complex64,
|
||||
XLA_FFI_DataType.C128: jnp.complex128,
|
||||
# XLA_FFI_DataType.TOKEN
|
||||
# XLA_FFI_DataType.F4E2M1FN: jnp.float4_e2m1fn.dtype,
|
||||
# XLA_FFI_DataType.F8E8M0FNU: jnp.float8_e8m0fnu.dtype,
|
||||
}
|
||||
|
||||
# newer types not supported by older versions
|
||||
if hasattr(jnp, "float8_e5m2"):
|
||||
constructors[XLA_FFI_DataType.F8E5M2] = jnp.float8_e5m2
|
||||
if hasattr(jnp, "float8_e3m4"):
|
||||
constructors[XLA_FFI_DataType.F8E3M4] = jnp.float8_e3m4
|
||||
if hasattr(jnp, "float8_e4m3"):
|
||||
constructors[XLA_FFI_DataType.F8E4M3] = jnp.float8_e4m3
|
||||
if hasattr(jnp, "float8_e4m3fn"):
|
||||
constructors[XLA_FFI_DataType.F8E4M3FN] = jnp.float8_e4m3fn
|
||||
if hasattr(jnp, "float8_e4m3b11fnuz"):
|
||||
constructors[XLA_FFI_DataType.F8E4M3B11FNUZ] = jnp.float8_e4m3b11fnuz
|
||||
if hasattr(jnp, "float8_e5m2fnuz"):
|
||||
constructors[XLA_FFI_DataType.F8E5M2FNUZ] = jnp.float8_e5m2fnuz
|
||||
if hasattr(jnp, "float8_e4m3fnuz"):
|
||||
constructors[XLA_FFI_DataType.F8E4M3FNUZ] = jnp.float8_e4m3fnuz
|
||||
|
||||
_xla_data_type_to_constructor = constructors
|
||||
|
||||
return _xla_data_type_to_constructor
|
||||
|
||||
|
||||
def _jnp_dtype_for_xla(xla_dtype):
|
||||
jnp = _get_jnp()
|
||||
return jnp.dtype(_get_xla_data_type_to_constructor()[xla_dtype])
|
||||
|
||||
|
||||
#######################################################################
|
||||
# ctypes structures and enums for XLA's FFI API:
|
||||
@@ -470,45 +535,6 @@ class XLA_FFI_CallFrame(ctypes.Structure):
|
||||
)
|
||||
|
||||
|
||||
_xla_data_type_to_constructor = {
|
||||
# XLA_FFI_DataType.INVALID
|
||||
XLA_FFI_DataType.PRED: jnp.bool,
|
||||
XLA_FFI_DataType.S8: jnp.int8,
|
||||
XLA_FFI_DataType.S16: jnp.int16,
|
||||
XLA_FFI_DataType.S32: jnp.int32,
|
||||
XLA_FFI_DataType.S64: jnp.int64,
|
||||
XLA_FFI_DataType.U8: jnp.uint8,
|
||||
XLA_FFI_DataType.U16: jnp.uint16,
|
||||
XLA_FFI_DataType.U32: jnp.uint32,
|
||||
XLA_FFI_DataType.U64: jnp.uint64,
|
||||
XLA_FFI_DataType.F16: jnp.float16,
|
||||
XLA_FFI_DataType.F32: jnp.float32,
|
||||
XLA_FFI_DataType.F64: jnp.float64,
|
||||
XLA_FFI_DataType.BF16: jnp.bfloat16,
|
||||
XLA_FFI_DataType.C64: jnp.complex64,
|
||||
XLA_FFI_DataType.C128: jnp.complex128,
|
||||
# XLA_FFI_DataType.TOKEN
|
||||
# XLA_FFI_DataType.F4E2M1FN: jnp.float4_e2m1fn.dtype,
|
||||
# XLA_FFI_DataType.F8E8M0FNU: jnp.float8_e8m0fnu.dtype,
|
||||
}
|
||||
|
||||
# newer types not supported by older versions
|
||||
if hasattr(jnp, "float8_e5m2"):
|
||||
_xla_data_type_to_constructor[XLA_FFI_DataType.F8E5M2] = jnp.float8_e5m2
|
||||
if hasattr(jnp, "float8_e3m4"):
|
||||
_xla_data_type_to_constructor[XLA_FFI_DataType.F8E3M4] = jnp.float8_e3m4
|
||||
if hasattr(jnp, "float8_e4m3"):
|
||||
_xla_data_type_to_constructor[XLA_FFI_DataType.F8E4M3] = jnp.float8_e4m3
|
||||
if hasattr(jnp, "float8_e4m3fn"):
|
||||
_xla_data_type_to_constructor[XLA_FFI_DataType.F8E4M3FN] = jnp.float8_e4m3fn
|
||||
if hasattr(jnp, "float8_e4m3b11fnuz"):
|
||||
_xla_data_type_to_constructor[XLA_FFI_DataType.F8E4M3B11FNUZ] = jnp.float8_e4m3b11fnuz
|
||||
if hasattr(jnp, "float8_e5m2fnuz"):
|
||||
_xla_data_type_to_constructor[XLA_FFI_DataType.F8E5M2FNUZ] = jnp.float8_e5m2fnuz
|
||||
if hasattr(jnp, "float8_e4m3fnuz"):
|
||||
_xla_data_type_to_constructor[XLA_FFI_DataType.F8E4M3FNUZ] = jnp.float8_e4m3fnuz
|
||||
|
||||
|
||||
########################################################################
|
||||
# Helpers for translating between ctypes and python types
|
||||
#######################################################################
|
||||
@@ -522,14 +548,14 @@ def decode_bytespan(span: XLA_FFI_ByteSpan):
|
||||
|
||||
def decode_scalar(scalar: XLA_FFI_Scalar):
|
||||
# TODO validate if dtype supported
|
||||
dtype = jnp.dtype(_xla_data_type_to_constructor[scalar.dtype])
|
||||
dtype = _jnp_dtype_for_xla(scalar.dtype)
|
||||
bytes = ctypes.string_at(scalar.value, dtype.itemsize)
|
||||
return np.frombuffer(bytes, dtype=dtype).reshape(())
|
||||
|
||||
|
||||
def decode_array(array: XLA_FFI_Array):
|
||||
# TODO validate if dtype supported
|
||||
dtype = jnp.dtype(_xla_data_type_to_constructor[array.dtype])
|
||||
dtype = _jnp_dtype_for_xla(array.dtype)
|
||||
bytes = ctypes.string_at(array.data, dtype.itemsize * array.size)
|
||||
return np.frombuffer(bytes, dtype=dtype)
|
||||
|
||||
@@ -612,7 +638,7 @@ def dtype_from_ffi(ffi_dtype):
|
||||
|
||||
|
||||
def jax_dtype_from_ffi(ffi_dtype):
|
||||
return _xla_data_type_to_constructor.get(ffi_dtype)
|
||||
return _get_xla_data_type_to_constructor().get(ffi_dtype)
|
||||
|
||||
|
||||
# Execution context (stream, stage)
|
||||
@@ -632,7 +658,7 @@ class FfiBuffer:
|
||||
|
||||
def __init__(self, xla_buffer):
|
||||
# TODO check if valid
|
||||
self.dtype = jnp.dtype(_xla_data_type_to_constructor[xla_buffer.dtype])
|
||||
self.dtype = _jnp_dtype_for_xla(xla_buffer.dtype)
|
||||
self.shape = tuple(xla_buffer.dims[i] for i in range(xla_buffer.rank))
|
||||
self.data = xla_buffer.data
|
||||
|
||||
@@ -1,2 +0,0 @@
|
||||
# SPDX-FileCopyrightText: Copyright (c) 2025 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
@@ -14,19 +14,17 @@
|
||||
# ==============================================================================
|
||||
|
||||
"""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}
|
||||
@@ -50,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 _refit_bvh_shim(
|
||||
# Model
|
||||
|
||||
@@ -14,17 +14,14 @@
|
||||
# ==============================================================================
|
||||
|
||||
"""DO NOT EDIT. This file is auto-generated."""
|
||||
|
||||
import dataclasses
|
||||
import functools
|
||||
|
||||
import jax
|
||||
import warp as wp
|
||||
|
||||
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
|
||||
from mujoco.mjx.warp import ffi
|
||||
import warp as wp
|
||||
|
||||
_m = mjwarp.Model(
|
||||
**{f.name: None for f in dataclasses.fields(mjwarp.Model) if f.init}
|
||||
@@ -48,7 +45,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 _collision_shim(
|
||||
# Model
|
||||
|
||||
@@ -22,10 +22,11 @@ from typing import Any, Callable, Optional, Sequence, Tuple, Union
|
||||
|
||||
import jax
|
||||
from jax import numpy as jp
|
||||
from mujoco.mjx.warp import types as mjx_warp_types
|
||||
import numpy as np
|
||||
import warp as wp
|
||||
from mujoco.mjx.third_party.warp._src.jax_experimental import ffi
|
||||
from mujoco.mjx.third_party.warp._src.jax import ffi as warp_ffi
|
||||
|
||||
from mujoco.mjx.warp import types as mjx_warp_types
|
||||
|
||||
|
||||
def flatten_signature(signature: inspect.Signature, args: Tuple[Any, ...]):
|
||||
@@ -97,7 +98,7 @@ def flatten_signature(signature: inspect.Signature, args: Tuple[Any, ...]):
|
||||
def jax_callable_variadic_tuple(
|
||||
func: Callable, # pylint: disable=g-bare-generic
|
||||
num_outputs: int = 1,
|
||||
graph_mode: ffi.GraphMode = ffi.GraphMode.WARP,
|
||||
graph_mode: warp_ffi.JaxCallableGraphMode = warp_ffi.JaxCallableGraphMode.WARP,
|
||||
vmap_method: Optional[str] = None,
|
||||
output_dims: Optional[dict[str, tuple[int, ...]]] = None,
|
||||
in_out_argnames: Optional[Sequence[str]] = None,
|
||||
@@ -126,7 +127,7 @@ def jax_callable_variadic_tuple(
|
||||
if new_signature.return_annotation is not inspect.Signature.empty:
|
||||
func_wrapper.__annotations__['return'] = new_signature.return_annotation
|
||||
|
||||
my_callable = ffi.jax_callable(
|
||||
my_callable = warp_ffi.jax_callable(
|
||||
func_wrapper,
|
||||
num_outputs=num_outputs,
|
||||
graph_mode=graph_mode,
|
||||
|
||||
@@ -23,6 +23,7 @@ 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}
|
||||
)
|
||||
@@ -45,7 +46,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 _forward_shim(
|
||||
# Model
|
||||
@@ -1609,8 +1609,8 @@ def _forward_jax_impl(m: types.Model, d: types.Data):
|
||||
m._impl.M_mulm_col,
|
||||
m._impl.M_mulm_madr,
|
||||
m._impl.M_mulm_rowadr,
|
||||
m._impl.M_rowadr,
|
||||
m._impl.M_rownnz,
|
||||
m.M_rowadr,
|
||||
m.M_rownnz,
|
||||
m._impl.M_tiles,
|
||||
m.actuator_acc0,
|
||||
m.actuator_actadr,
|
||||
@@ -3899,8 +3899,8 @@ def _step_jax_impl(m: types.Model, d: types.Data):
|
||||
m._impl.M_mulm_col,
|
||||
m._impl.M_mulm_madr,
|
||||
m._impl.M_mulm_rowadr,
|
||||
m._impl.M_rowadr,
|
||||
m._impl.M_rownnz,
|
||||
m.M_rowadr,
|
||||
m.M_rownnz,
|
||||
m._impl.M_tiles,
|
||||
m.actuator_acc0,
|
||||
m.actuator_actadr,
|
||||
|
||||
@@ -14,19 +14,17 @@
|
||||
# ==============================================================================
|
||||
|
||||
"""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}
|
||||
@@ -50,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,6 +14,7 @@
|
||||
# ==============================================================================
|
||||
import functools
|
||||
import os
|
||||
import tempfile
|
||||
|
||||
from absl.testing import absltest
|
||||
from absl.testing import parameterized
|
||||
@@ -70,11 +71,21 @@ def _get_model_data_rc(xml, batch_size, render_seg=False):
|
||||
|
||||
class RenderTest(parameterized.TestCase):
|
||||
|
||||
@classmethod
|
||||
def setUpClass(cls):
|
||||
super().setUpClass()
|
||||
if mjxw.WARP_INSTALLED:
|
||||
cls.tempdir = tempfile.TemporaryDirectory()
|
||||
wp.config.kernel_cache_dir = cls.tempdir.name
|
||||
|
||||
@classmethod
|
||||
def tearDownClass(cls):
|
||||
super().tearDownClass()
|
||||
if hasattr(cls, 'tempdir'):
|
||||
cls.tempdir.cleanup()
|
||||
|
||||
def setUp(self):
|
||||
super().setUp()
|
||||
if mjxw.WARP_INSTALLED:
|
||||
tempdir = '/tmp/wp_kernel_cache_dir_RenderTest'
|
||||
wp.config.kernel_cache_dir = tempdir
|
||||
np.random.seed(0)
|
||||
|
||||
def _maybe_skip(self):
|
||||
@@ -251,11 +262,18 @@ class RenderTest(parameterized.TestCase):
|
||||
class RenderContextGarbageCollectionTest(absltest.TestCase):
|
||||
"""Tests that RenderContext cleans up buffers on deletion."""
|
||||
|
||||
def setUp(self):
|
||||
super().setUp()
|
||||
@classmethod
|
||||
def setUpClass(cls):
|
||||
super().setUpClass()
|
||||
if mjxw.WARP_INSTALLED:
|
||||
tempdir = '/tmp/wp_kernel_cache_dir_RenderContextGCTest'
|
||||
wp.config.kernel_cache_dir = tempdir
|
||||
cls.tempdir = tempfile.TemporaryDirectory()
|
||||
wp.config.kernel_cache_dir = cls.tempdir.name
|
||||
|
||||
@classmethod
|
||||
def tearDownClass(cls):
|
||||
super().tearDownClass()
|
||||
if hasattr(cls, 'tempdir'):
|
||||
cls.tempdir.cleanup()
|
||||
|
||||
def _maybe_skip(self):
|
||||
if not mjxw.WARP_INSTALLED:
|
||||
|
||||
@@ -14,17 +14,14 @@
|
||||
# ==============================================================================
|
||||
|
||||
"""DO NOT EDIT. This file is auto-generated."""
|
||||
|
||||
import dataclasses
|
||||
import functools
|
||||
|
||||
import jax
|
||||
import warp as wp
|
||||
|
||||
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
|
||||
from mujoco.mjx.warp import ffi
|
||||
import warp as wp
|
||||
|
||||
_m = mjwarp.Model(
|
||||
**{f.name: None for f in dataclasses.fields(mjwarp.Model) if f.init}
|
||||
@@ -48,7 +45,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 _kinematics_shim(
|
||||
# Model
|
||||
|
||||
@@ -32,10 +32,10 @@ from mujoco.mjx.warp import collision_driver as wp_collision
|
||||
from mujoco.mjx.warp import forward as wp_forward
|
||||
from mujoco.mjx.warp import io
|
||||
from mujoco.mjx.warp import smooth as wp_smooth
|
||||
from mujoco.mjx.warp import types as mjxw_types
|
||||
import mujoco.mjx.third_party.mujoco_warp as mjwarp
|
||||
import numpy as np
|
||||
import warp as wp
|
||||
from mujoco.mjx.third_party.warp._src.jax_experimental import ffi as warp_ffi
|
||||
|
||||
|
||||
_MODELFILE = flags.DEFINE_string(
|
||||
@@ -280,7 +280,7 @@ def _main(_: Sequence[str]):
|
||||
m = mujoco.MjModel.from_xml_path(modelfile)
|
||||
|
||||
benchmark_type = _BENCHMARK.value
|
||||
graph_mode = getattr(warp_ffi.GraphMode, _GRAPH_MODE.value)
|
||||
graph_mode = getattr(mjxw_types.GraphMode, _GRAPH_MODE.value)
|
||||
|
||||
# Only allocate the model needed for the specific benchmark
|
||||
mx = None
|
||||
|
||||
@@ -23,7 +23,6 @@ from jax import tree_util
|
||||
from jax.interpreters import batching
|
||||
from mujoco.mjx._src import dataclasses as mjx_dataclasses
|
||||
import numpy as np
|
||||
|
||||
if typing.TYPE_CHECKING:
|
||||
GraphMode = int
|
||||
|
||||
@@ -33,7 +32,9 @@ if typing.TYPE_CHECKING:
|
||||
|
||||
else:
|
||||
try:
|
||||
from warp._src.jax_experimental.ffi import GraphMode
|
||||
from mujoco.mjx.third_party.warp._src.jax import ffi as warp_ffi
|
||||
|
||||
GraphMode = warp_ffi.JaxCallableGraphMode
|
||||
from mujoco.mjx.third_party.mujoco_warp._src import types as mjwp_types
|
||||
|
||||
Callback = mjwp_types.Callback
|
||||
@@ -170,7 +171,6 @@ class ModelWarp(PyTreeNode):
|
||||
D_diag: np.ndarray
|
||||
D_rowadr: np.ndarray
|
||||
D_rownnz: np.ndarray
|
||||
M_colind: np.ndarray
|
||||
M_elemid: np.ndarray
|
||||
M_fullm_i: np.ndarray
|
||||
M_fullm_j: np.ndarray
|
||||
@@ -180,8 +180,6 @@ class ModelWarp(PyTreeNode):
|
||||
M_mulm_col: np.ndarray
|
||||
M_mulm_madr: np.ndarray
|
||||
M_mulm_rowadr: np.ndarray
|
||||
M_rowadr: np.ndarray
|
||||
M_rownnz: np.ndarray
|
||||
M_tiles: Tuple[TileSet, ...]
|
||||
actuator_delay: np.ndarray
|
||||
actuator_history: np.ndarray
|
||||
|
||||
+1
-1
@@ -37,7 +37,7 @@ dependencies = [
|
||||
|
||||
[project.optional-dependencies]
|
||||
warp = [
|
||||
"warp-lang==1.13.0",
|
||||
"warp-lang==1.14.0",
|
||||
]
|
||||
dev = [
|
||||
"isort",
|
||||
|
||||
Reference in New Issue
Block a user