diff --git a/mjx/cuda_requirements.txt b/mjx/cuda_requirements.txt index a726944a..ec304a7f 100644 --- a/mjx/cuda_requirements.txt +++ b/mjx/cuda_requirements.txt @@ -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 diff --git a/mjx/mujoco/mjx/_src/io.py b/mjx/mujoco/mjx/_src/io.py index b896ccc7..fc1a6595 100644 --- a/mjx/mujoco/mjx/_src/io.py +++ b/mjx/mujoco/mjx/_src/io.py @@ -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. diff --git a/mjx/mujoco/mjx/codegen/generate_warp_types.py b/mjx/mujoco/mjx/codegen/generate_warp_types.py index 98b57215..6c9e3c05 100644 --- a/mjx/mujoco/mjx/codegen/generate_warp_types.py +++ b/mjx/mujoco/mjx/codegen/generate_warp_types.py @@ -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) diff --git a/mjx/mujoco/mjx/third_party/mujoco_warp/_src/render.py b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/render.py index 5859c6ee..e33dd623 100644 --- a/mjx/mujoco/mjx/third_party/mujoco_warp/_src/render.py +++ b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/render.py @@ -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, ) diff --git a/mjx/mujoco/mjx/third_party/warp/_src/jax/__init__.py b/mjx/mujoco/mjx/third_party/warp/_src/jax/__init__.py new file mode 100644 index 00000000..9f22de18 --- /dev/null +++ b/mjx/mujoco/mjx/third_party/warp/_src/jax/__init__.py @@ -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) diff --git a/mjx/mujoco/mjx/third_party/warp/_src/jax_experimental/custom_call.py b/mjx/mujoco/mjx/third_party/warp/_src/jax/custom_call.py similarity index 96% rename from mjx/mujoco/mjx/third_party/warp/_src/jax_experimental/custom_call.py rename to mjx/mujoco/mjx/third_party/warp/_src/jax/custom_call.py index 6f63ba66..cb2fce91 100644 --- a/mjx/mujoco/mjx/third_party/warp/_src/jax_experimental/custom_call.py +++ b/mjx/mujoco/mjx/third_party/warp/_src/jax/custom_call.py @@ -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(";") diff --git a/mjx/mujoco/mjx/third_party/warp/_src/jax_experimental/ffi.py b/mjx/mujoco/mjx/third_party/warp/_src/jax/ffi.py similarity index 91% rename from mjx/mujoco/mjx/third_party/warp/_src/jax_experimental/ffi.py rename to mjx/mujoco/mjx/third_party/warp/_src/jax/ffi.py index ac7b0f26..8a1d6cad 100644 --- a/mjx/mujoco/mjx/third_party/warp/_src/jax_experimental/ffi.py +++ b/mjx/mujoco/mjx/third_party/warp/_src/jax/ffi.py @@ -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,) diff --git a/mjx/mujoco/mjx/third_party/warp/_src/jax_experimental/xla_ffi.py b/mjx/mujoco/mjx/third_party/warp/_src/jax/xla_ffi.py similarity index 86% rename from mjx/mujoco/mjx/third_party/warp/_src/jax_experimental/xla_ffi.py rename to mjx/mujoco/mjx/third_party/warp/_src/jax/xla_ffi.py index 911ca2f1..a032c5b9 100644 --- a/mjx/mujoco/mjx/third_party/warp/_src/jax_experimental/xla_ffi.py +++ b/mjx/mujoco/mjx/third_party/warp/_src/jax/xla_ffi.py @@ -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 diff --git a/mjx/mujoco/mjx/third_party/warp/_src/jax_experimental/__init__.py b/mjx/mujoco/mjx/third_party/warp/_src/jax_experimental/__init__.py deleted file mode 100644 index 1a8431c3..00000000 --- a/mjx/mujoco/mjx/third_party/warp/_src/jax_experimental/__init__.py +++ /dev/null @@ -1,2 +0,0 @@ -# SPDX-FileCopyrightText: Copyright (c) 2025 NVIDIA CORPORATION & AFFILIATES. All rights reserved. -# SPDX-License-Identifier: Apache-2.0 diff --git a/mjx/mujoco/mjx/warp/bvh.py b/mjx/mujoco/mjx/warp/bvh.py index 1301aa88..7f648694 100644 --- a/mjx/mujoco/mjx/warp/bvh.py +++ b/mjx/mujoco/mjx/warp/bvh.py @@ -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 diff --git a/mjx/mujoco/mjx/warp/collision_driver.py b/mjx/mujoco/mjx/warp/collision_driver.py index 1de26245..bb020a16 100644 --- a/mjx/mujoco/mjx/warp/collision_driver.py +++ b/mjx/mujoco/mjx/warp/collision_driver.py @@ -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 diff --git a/mjx/mujoco/mjx/warp/ffi.py b/mjx/mujoco/mjx/warp/ffi.py index ce46b599..9e956991 100644 --- a/mjx/mujoco/mjx/warp/ffi.py +++ b/mjx/mujoco/mjx/warp/ffi.py @@ -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, diff --git a/mjx/mujoco/mjx/warp/forward.py b/mjx/mujoco/mjx/warp/forward.py index c3238c02..c9e386ea 100644 --- a/mjx/mujoco/mjx/warp/forward.py +++ b/mjx/mujoco/mjx/warp/forward.py @@ -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, diff --git a/mjx/mujoco/mjx/warp/render.py b/mjx/mujoco/mjx/warp/render.py index bd60fcf7..5ab3ae69 100644 --- a/mjx/mujoco/mjx/warp/render.py +++ b/mjx/mujoco/mjx/warp/render.py @@ -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 diff --git a/mjx/mujoco/mjx/warp/render_test.py b/mjx/mujoco/mjx/warp/render_test.py index 280e846d..b6c1c048 100644 --- a/mjx/mujoco/mjx/warp/render_test.py +++ b/mjx/mujoco/mjx/warp/render_test.py @@ -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: diff --git a/mjx/mujoco/mjx/warp/smooth.py b/mjx/mujoco/mjx/warp/smooth.py index 862b6d07..bf26a123 100644 --- a/mjx/mujoco/mjx/warp/smooth.py +++ b/mjx/mujoco/mjx/warp/smooth.py @@ -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 diff --git a/mjx/mujoco/mjx/warp/testspeed.py b/mjx/mujoco/mjx/warp/testspeed.py index d8869263..296e75c8 100644 --- a/mjx/mujoco/mjx/warp/testspeed.py +++ b/mjx/mujoco/mjx/warp/testspeed.py @@ -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 diff --git a/mjx/mujoco/mjx/warp/types.py b/mjx/mujoco/mjx/warp/types.py index 6e953948..3474b32c 100644 --- a/mjx/mujoco/mjx/warp/types.py +++ b/mjx/mujoco/mjx/warp/types.py @@ -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 diff --git a/mjx/pyproject.toml b/mjx/pyproject.toml index f8a1d3da..88d04d77 100644 --- a/mjx/pyproject.toml +++ b/mjx/pyproject.toml @@ -37,7 +37,7 @@ dependencies = [ [project.optional-dependencies] warp = [ - "warp-lang==1.13.0", + "warp-lang==1.14.0", ] dev = [ "isort",