diff --git a/mjx/cuda_requirements.txt b/mjx/cuda_requirements.txt index f51f641b..d6bbdc7a 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.15.0 \ - --hash=sha256:06ef42c3ce522749268056653ea253645eb4a96ca164af74ae5113f0d55a6970 \ - --hash=sha256:3c0711b39ad98ff924d29654b028ec4eaa02330e207ae3efc72445e9cde05b34 \ - --hash=sha256:95c169f28bd7d6c78ac4ad62e2df1e61a096033748f757157fa4551aed80d010 \ - --hash=sha256:ee06e2e76bff84088a09e64315e71ed0e55e94c936ec664d7b05ebd56f66a363 +warp-lang==1.16.0 \ + --hash=sha256:19a7590c4484e8250ab19165311d2a1774138663b361c41e8be2865c7515c4ae \ + --hash=sha256:8690c9096e0a271985339d4aa4b37675e7d62852620ca861ac032dc85b124542 \ + --hash=sha256:96449fc1e3b354185e2f09434fb794b5953ab2e8673b104d7e7f48d5d418bb35 \ + --hash=sha256:d4715171bd6821436b82293d774c9a082ace08215c189572b4348d1e48fb8ee8 diff --git a/mjx/mujoco/mjx/_src/io.py b/mjx/mujoco/mjx/_src/io.py index 54bf16c7..fe3d4704 100644 --- a/mjx/mujoco/mjx/_src/io.py +++ b/mjx/mujoco/mjx/_src/io.py @@ -568,7 +568,11 @@ def put_model( return _put_model_jax(m, device) elif impl == types.Impl.WARP: _check_warp_installed() - graph_mode = graph_mode or getattr(mjxw.types.GraphMode, 'WARP') + if graph_mode is None: + if device.platform == 'cpu': + graph_mode = getattr(mjxw.types.GraphMode, 'JAX') + else: + graph_mode = getattr(mjxw.types.GraphMode, 'WARP') return _put_model_warp(m, graph_mode, device, batch_sizes=batch_sizes) elif impl == types.Impl.CPP: return _put_model_cpp(m, device, keepalive_refs=keepalive_refs) diff --git a/mjx/mujoco/mjx/third_party/warp/_src/jax/__init__.py b/mjx/mujoco/mjx/third_party/warp/_src/jax/__init__.py index 9f22de18..484ad0d8 100644 --- a/mjx/mujoco/mjx/third_party/warp/_src/jax/__init__.py +++ b/mjx/mujoco/mjx/third_party/warp/_src/jax/__init__.py @@ -3,8 +3,6 @@ import warp -_wp_module_name_ = "warp.jax" - def device_to_jax(warp_device: warp.DeviceLike): """Return the Jax device corresponding to a Warp device. diff --git a/mjx/mujoco/mjx/third_party/warp/_src/jax/custom_call.py b/mjx/mujoco/mjx/third_party/warp/_src/jax/custom_call.py index 416905e6..153e5305 100644 --- a/mjx/mujoco/mjx/third_party/warp/_src/jax/custom_call.py +++ b/mjx/mujoco/mjx/third_party/warp/_src/jax/custom_call.py @@ -10,8 +10,6 @@ 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.custom_call" - _jax_warp_p = None # Holder for the custom callback to keep it alive. diff --git a/mjx/mujoco/mjx/third_party/warp/_src/jax/ffi.py b/mjx/mujoco/mjx/third_party/warp/_src/jax/ffi.py index ef7b0f43..0caa291a 100644 --- a/mjx/mujoco/mjx/third_party/warp/_src/jax/ffi.py +++ b/mjx/mujoco/mjx/third_party/warp/_src/jax/ffi.py @@ -6,6 +6,7 @@ from __future__ import annotations import collections import ctypes import inspect +import operator import threading import traceback from collections.abc import Callable @@ -13,7 +14,13 @@ from enum import IntEnum import warp as wp from warp._src.codegen import get_full_arg_spec, make_full_qualified_name -from warp._src.context import CudaMemcpyKind, _build_launch_bounds +from warp._src.context import ( + CudaMemcpyKind, + _build_launch_bounds, + _raise_cuda_launch_error, + _validate_cluster_launch, + invoke, +) from mujoco.mjx.third_party.warp._src.jax import get_jax_device from warp._src.logger import log_warning from warp._src.types import ( @@ -26,8 +33,6 @@ from warp._src.types import ( from .xla_ffi import * -_wp_module_name_ = "warp.jax.ffi" - # Holders for the custom callbacks to keep them alive. _FFI_KERNEL_REGISTRY: dict[tuple, FfiKernel] = {} _FFI_DIFF_KERNEL_REGISTRY: dict[tuple, Callable] = {} @@ -63,6 +68,31 @@ def check_jax_version(): raise RuntimeError(msg) +# JAX platform identifiers are case-sensitive and intentionally use different casing. +_FFI_PLATFORM_CPU = "cpu" +_FFI_PLATFORM_CUDA = "CUDA" + + +def _register_ffi_targets(name, callback): + """Register one FFI callback for CPU and CUDA without initializing either backend.""" + jax = _get_jax() + callback_type = ctypes.CFUNCTYPE(ctypes.c_void_p, ctypes.POINTER(XLA_FFI_CallFrame)) + callback_funcs = {} + + for platform in (_FFI_PLATFORM_CPU, _FFI_PLATFORM_CUDA): + + def platform_callback(call_frame, platform=platform): + return callback(call_frame, platform) + + callback_func = callback_type(platform_callback) + callback_address = ctypes.cast(callback_func, ctypes.c_void_p) + callback_capsule = jax.ffi.pycapsule(callback_address.value) + jax.ffi.register_ffi_target(name, callback_capsule, platform=platform) + callback_funcs[platform] = callback_func + + return callback_funcs + + def collapse_batch_dims(shape, desired_ndim): # roll leading batch dims into one while len(shape) > desired_ndim: @@ -111,6 +141,65 @@ class JaxModulePreloadMode(IntEnum): ModulePreloadMode = JaxModulePreloadMode +def _get_ffi_block_dim(device, block_dim=None): + """Resolve the FFI block dimension for ``device``.""" + if device.is_cpu: + # Remove this override if CPU launches gain configurable block dimensions. + return 1 + return 256 if block_dim is None else block_dim + + +def _load_ffi_module(module, device, block_dim=None): + module_exec = module.load(device, _get_ffi_block_dim(device, block_dim)) + if module_exec is None: + raise RuntimeError( + f"Failed to load Warp module '{module.name}' on device '{device}' after a previous build failure" + ) + return module_exec + + +def _preload_ffi_module(module, mode, block_dim=None): + if mode == JaxModulePreloadMode.NONE: + return + + if mode == JaxModulePreloadMode.CURRENT_DEVICE: + jax_device = get_jax_device() + try: + device = wp.device_from_jax(jax_device) + except (IndexError, RuntimeError): + return + _load_ffi_module(module, device, block_dim) + return + + if mode == JaxModulePreloadMode.ALL_DEVICES: + jax = _get_jax() + devices = [] + mapped_device_ids = set() + for backend in ("cpu", "cuda"): + try: + jax_devices = jax.local_devices(backend=backend) + except RuntimeError: + continue + + for jax_device in jax_devices: + try: + device = wp.device_from_jax(jax_device) + except (IndexError, RuntimeError): + continue + + device_id = id(device) + if device_id in mapped_device_ids: + continue + mapped_device_ids.add(device_id) + devices.append(device) + + for device in devices: + _load_ffi_module(module, device, block_dim) + return + + raise ValueError(f"Unsupported JAX module preload mode '{mode}'") + + class FfiArg: def __init__(self, name, type, in_out=False): self.name = name @@ -155,6 +244,7 @@ class FfiKernel: num_outputs, vmap_method, launch_dims, + block_dim, output_dims, in_out_argnames, module_preload_mode, @@ -165,6 +255,7 @@ class FfiKernel: self.num_outputs = num_outputs self.vmap_method = vmap_method self.launch_dims = launch_dims + self.block_dim = block_dim self.output_dims = output_dims self.module_preload_mode = module_preload_mode self.has_side_effect = has_side_effect @@ -228,20 +319,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) - ) - - self.callback_func_cuda = FFI_CCALLFUNC(lambda call_frame: self.ffi_callback(call_frame, platform="CUDA")) - ffi_ccall_address_cuda = ctypes.cast(self.callback_func_cuda, ctypes.c_void_p) - ffi_capsule_cuda = jax.ffi.pycapsule(ffi_ccall_address_cuda.value) - jax.ffi.register_ffi_target(self.name, ffi_capsule_cuda, platform="CUDA") - - self.callback_func_host = FFI_CCALLFUNC(lambda call_frame: self.ffi_callback(call_frame, platform="Host")) - ffi_ccall_address_host = ctypes.cast(self.callback_func_host, ctypes.c_void_p) - ffi_capsule_host = jax.ffi.pycapsule(ffi_ccall_address_host.value) - jax.ffi.register_ffi_target(self.name, ffi_capsule_host, platform="Host") + self.callback_funcs = _register_ffi_targets(self.name, self.ffi_callback) def __call__(self, *args, output_dims=None, launch_dims=None, vmap_method=None): jax = _get_jax() @@ -332,20 +410,7 @@ class FfiKernel: has_side_effect=self.has_side_effect, ) - # preload on the specified devices - 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 == 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 + _preload_ffi_module(self.kernel.module, self.module_preload_mode, self.block_dim) # save launch data to be retrieved by callback launch_id = self.launch_id @@ -356,7 +421,7 @@ class FfiKernel: return call(*args, launch_id=launch_id) - def ffi_callback(self, call_frame, platform="CUDA"): + def ffi_callback(self, call_frame, platform): try: # On the first call, XLA runtime will query the API version and traits # metadata using the |extension| field. Let us respond to that query @@ -368,8 +433,8 @@ class FfiKernel: metadata_ext = ctypes.cast(extension, ctypes.POINTER(XLA_FFI_Metadata_Extension)) 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 platform == "CUDA": + if platform == _FFI_PLATFORM_CUDA: + # Turn on CUDA graphs for this handler. metadata_ext.contents.metadata.contents.traits = ( XLA_FFI_Handler_TraitsBits.COMMAND_BUFFER_COMPATIBLE ) @@ -392,8 +457,6 @@ class FfiKernel: assert num_inputs == self.num_inputs assert num_outputs == self.num_outputs - # first kernel param is the launch bounds - kernel_params = (ctypes.c_void_p * (1 + self.num_kernel_args))() arg_refs = [] batch_size = None @@ -411,13 +474,11 @@ class FfiKernel: ) strides = strides_from_shape(shape, input_arg.type.dtype) arg = array_t(buffer.data, 0, input_arg.type.ndim, shape, strides) - kernel_params[i + 1] = ctypes.addressof(arg) arg_refs.append(arg) # keep a reference else: # scalar argument, get stashed value value = launch_desc.static_inputs[input_arg.name] arg = input_arg.type._type_(value) - kernel_params[i + 1] = ctypes.addressof(arg) arg_refs.append(arg) # keep a reference # pure output args (skip in-out FFI buffers) @@ -433,9 +494,40 @@ class FfiKernel: ) strides = strides_from_shape(shape, output_arg.type.dtype) arg = array_t(buffer.data, 0, output_arg.type.ndim, shape, strides) - kernel_params[num_inputs + i + 1] = ctypes.addressof(arg) arg_refs.append(arg) # keep a reference + if platform == _FFI_PLATFORM_CPU: + if not wp.is_cpu_available(): + return create_ffi_error( + call_frame.contents.api, + XLA_FFI_Error_Code.FAILED_PRECONDITION, + "This Warp build does not include CPU support", + ) + device = wp.get_device("cpu") + stream = None + elif platform == _FFI_PLATFORM_CUDA: + if wp._src.context.runtime is None: + wp.init() + if not wp._src.context.runtime.is_cuda_enabled: + return create_ffi_error( + call_frame.contents.api, + XLA_FFI_Error_Code.FAILED_PRECONDITION, + "This Warp build does not include CUDA support", + ) + + device = wp.get_cuda_device(get_device_ordinal_from_callframe(call_frame.contents)) + stream = get_stream_from_callframe(call_frame.contents) + else: + return create_invalid_argument_ffi_error( + call_frame.contents.api, + f"Unsupported JAX FFI platform '{platform}'", + ) + + # Preloading is best-effort, so the callback's actual device must + # load the module before reading code-generation metadata. + module_exec = _load_ffi_module(self.kernel.module, device, self.block_dim) + block_dim = module_exec.block_dim + # determine launch bounds if launch_desc.launch_dims is None: # infer launch dims from argument shape, works with vmap @@ -449,42 +541,43 @@ class FfiKernel: launch_dims = (batch_size * launch_dims[0], *launch_dims[1:]) launch_bounds = _build_launch_bounds(launch_dims, self.kernel.adj.kernel_dim) - kernel_params[0] = ctypes.addressof(launch_bounds) - # get device and stream - if platform == "CUDA": - device = wp.get_cuda_device(get_device_ordinal_from_callframe(call_frame.contents)) - stream = get_stream_from_callframe(call_frame.contents) - else: - device = wp.get_device("cpu") - stream = None + if platform == _FFI_PLATFORM_CPU: + hooks = module_exec.get_kernel_hooks(self.kernel) + if hooks.forward is None: + raise RuntimeError("Failed to find CPU kernel entry point") + invoke(self.kernel, hooks, [launch_bounds, *arg_refs], adjoint=False) + return None + + kernel_params = (ctypes.c_void_p * (1 + self.num_kernel_args))( + ctypes.addressof(launch_bounds), + *(ctypes.addressof(arg) for arg in arg_refs), + ) # get kernel hooks - hooks = self.kernel.module.get_kernel_hooks(self.kernel, device) - assert hooks.forward, "Failed to find kernel entry point" + hooks = module_exec.get_kernel_hooks(self.kernel) + if hooks.forward is None: + raise RuntimeError("Failed to find CUDA kernel entry point") - # launch the kernel - if device.is_cuda: - wp._src.context.runtime.core.wp_cuda_launch_kernel( - device.context, - hooks.forward, - launch_bounds.size, - 0, - 256, - int(self.kernel.grid_stride), - hooks.cluster_dim, - hooks.forward_smem_bytes, - kernel_params, - stream, - None, # apic_info - ) - else: - wp._src.context.runtime.core.wp_cpu_launch_kernel( - device.context, - hooks.forward, - launch_bounds.size, - kernel_params, - ) + # reject non-cluster-aligned grids with a clear Python error instead + # of a cryptic native CUDA error (configured block_dim, max_blocks=0 below) + _validate_cluster_launch(hooks.cluster_dim, launch_bounds.size, block_dim, 0) + + # launch the kernel (cluster_dim is cached on the hooks at load time) + if wp._src.context.runtime.core.wp_cuda_launch_kernel( + device.context, + hooks.forward, + launch_bounds.size, + 0, + block_dim, + int(self.kernel.grid_stride), + hooks.cluster_dim, + hooks.forward_smem_bytes, + kernel_params, + stream, + None, # apic_info + ): + _raise_cuda_launch_error(self.kernel, device) except Exception as e: print(traceback.format_exc()) @@ -637,20 +730,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) - ) - - self.callback_func_cuda = FFI_CCALLFUNC(lambda call_frame: self.ffi_callback(call_frame, platform="CUDA")) - ffi_ccall_address_cuda = ctypes.cast(self.callback_func_cuda, ctypes.c_void_p) - ffi_capsule_cuda = jax.ffi.pycapsule(ffi_ccall_address_cuda.value) - jax.ffi.register_ffi_target(self.name, ffi_capsule_cuda, platform="CUDA") - - self.callback_func_host = FFI_CCALLFUNC(lambda call_frame: self.ffi_callback(call_frame, platform="Host")) - ffi_ccall_address_host = ctypes.cast(self.callback_func_host, ctypes.c_void_p) - ffi_capsule_host = jax.ffi.pycapsule(ffi_ccall_address_host.value) - jax.ffi.register_ffi_target(self.name, ffi_capsule_host, platform="Host") + self.callback_funcs = _register_ffi_targets(self.name, self.ffi_callback) def __call__(self, *args, output_dims=None, vmap_method=None): jax = _get_jax() @@ -729,21 +809,9 @@ class FfiCallable: has_side_effect=self.has_side_effect, ) - # 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 == JaxModulePreloadMode.CURRENT_DEVICE: - device = wp.device_from_jax(get_jax_device()) - module.load(device) - 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 + _preload_ffi_module(module, self.module_preload_mode) # save call data to be retrieved by callback call_id = self.call_id @@ -751,7 +819,25 @@ class FfiCallable: self.call_id += 1 return call(*args, call_id=call_id) - def ffi_callback(self, call_frame, platform="CUDA"): + def _build_arg_list(self, inputs, outputs, call_desc, device): + arg_list = [] + + for i, arg in enumerate(self.input_args): + if arg.is_array: + buffer = inputs[i].contents + shape = collapse_batch_dims(buffer.dims[: buffer.rank - arg.dtype_ndim], arg.type.ndim) + arg_list.append(wp.array(ptr=buffer.data, dtype=arg.type.dtype, shape=shape, device=device)) + else: + arg_list.append(call_desc.static_inputs[arg.name]) + + for i, arg in enumerate(self.output_args): + buffer = outputs[i + self.num_in_out].contents + shape = collapse_batch_dims(buffer.dims[: buffer.rank - arg.dtype_ndim], arg.type.ndim) + arg_list.append(wp.array(ptr=buffer.data, dtype=arg.type.dtype, shape=shape, device=device)) + + return arg_list + + def ffi_callback(self, call_frame, platform): try: # On the first call, XLA runtime will query the API version and traits # metadata using the |extension| field. Let us respond to that query @@ -763,8 +849,8 @@ class FfiCallable: metadata_ext = ctypes.cast(extension, ctypes.POINTER(XLA_FFI_Metadata_Extension)) 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 JaxCallableGraphMode.JAX and platform == "CUDA": + # Turn on CUDA graphs for this handler. + if platform == _FFI_PLATFORM_CUDA and self.graph_mode is JaxCallableGraphMode.JAX: metadata_ext.contents.metadata.contents.traits = ( XLA_FFI_Handler_TraitsBits.COMMAND_BUFFER_COMPATIBLE ) @@ -791,37 +877,45 @@ class FfiCallable: assert num_inputs == self.num_inputs assert num_outputs == self.num_outputs - if platform == "Host": + if platform == _FFI_PLATFORM_CPU: + if self.graph_mode not in (JaxCallableGraphMode.NONE, JaxCallableGraphMode.JAX): + return create_invalid_argument_ffi_error( + call_frame.contents.api, + f"JaxCallableGraphMode.{self.graph_mode.name} is not supported for JAX CPU FFI calls", + ) + + if not wp.is_cpu_available(): + return create_ffi_error( + call_frame.contents.api, + XLA_FFI_Error_Code.FAILED_PRECONDITION, + "This Warp build does not include CPU support", + ) device = wp.get_device("cpu") - # reconstruct the argument list - arg_list = [] - - # input and in-out args - for i, arg in enumerate(self.input_args): - if arg.is_array: - buffer = inputs[i].contents - shape = collapse_batch_dims(buffer.dims[: buffer.rank - arg.dtype_ndim], arg.type.ndim) - arr = wp.array(ptr=buffer.data, dtype=arg.type.dtype, shape=shape, device=device) - arg_list.append(arr) - else: - # scalar argument, get stashed value - value = call_desc.static_inputs[arg.name] - arg_list.append(value) - - # pure output args (skip in-out FFI buffers) - for i, arg in enumerate(self.output_args): - buffer = outputs[i + self.num_in_out].contents - shape = collapse_batch_dims(buffer.dims[: buffer.rank - arg.dtype_ndim], arg.type.ndim) - arr = wp.array(ptr=buffer.data, dtype=arg.type.dtype, shape=shape, device=device) - arg_list.append(arr) - - # call the Python function with reconstructed arguments + module = wp.get_module(self.func.__module__) + _load_ffi_module(module, device) + arg_list = self._build_arg_list(inputs, outputs, call_desc, device) with wp.ScopedDevice(device): self.func(*arg_list) - return + return None + + if platform != _FFI_PLATFORM_CUDA: + return create_invalid_argument_ffi_error( + call_frame.contents.api, + f"Unsupported JAX FFI platform '{platform}'", + ) + + if wp._src.context.runtime is None: + wp.init() + if not wp._src.context.runtime.is_cuda_enabled: + return create_ffi_error( + call_frame.contents.api, + XLA_FFI_Error_Code.FAILED_PRECONDITION, + "This Warp build does not include CUDA support", + ) cuda_stream = get_stream_from_callframe(call_frame.contents) device_ordinal = get_device_ordinal_from_callframe(call_frame.contents) + device = wp.get_cuda_device(device_ordinal) if self.graph_mode == JaxCallableGraphMode.WARP: # check if we already captured an identical call @@ -925,35 +1019,14 @@ class FfiCallable: # early out return - device_ordinal = get_device_ordinal_from_callframe(call_frame.contents) - device = wp.get_cuda_device(device_ordinal) + _load_ffi_module(wp.get_module(self.func.__module__), device) stream = wp.Stream(device, cuda_stream=cuda_stream) - # reconstruct the argument list - arg_list = [] - - # input and in-out args - for i, arg in enumerate(self.input_args): - if arg.is_array: - buffer = inputs[i].contents - shape = collapse_batch_dims(buffer.dims[: buffer.rank - arg.dtype_ndim], arg.type.ndim) - arr = wp.array(ptr=buffer.data, dtype=arg.type.dtype, shape=shape, device=device) - arg_list.append(arr) - else: - # scalar argument, get stashed value - value = call_desc.static_inputs[arg.name] - arg_list.append(value) - - # pure output args (skip in-out FFI buffers) - for i, arg in enumerate(self.output_args): - buffer = outputs[i + self.num_in_out].contents - shape = collapse_batch_dims(buffer.dims[: buffer.rank - arg.dtype_ndim], arg.type.ndim) - arr = wp.array(ptr=buffer.data, dtype=arg.type.dtype, shape=shape, device=device) - arg_list.append(arr) + arg_list = self._build_arg_list(inputs, outputs, call_desc, device) # call the Python function with reconstructed arguments - with wp.ScopedStream(stream, sync_enter=False) if stream else wp.ScopedDevice(device): - if stream and stream.is_capturing: + with wp.ScopedStream(stream, sync_enter=False): + if stream.is_capturing: # capturing with JAX with wp.ScopedCapture(external=True) as capture: self.func(*arg_list) @@ -961,7 +1034,7 @@ class FfiCallable: # keep a reference to the capture object to prevent required modules getting unloaded call_desc.capture = capture - elif self.graph_mode == JaxCallableGraphMode.WARP and device.is_cuda: + elif self.graph_mode == JaxCallableGraphMode.WARP: # capturing with WARP with wp.ScopedCapture() as capture: self.func(*arg_list) @@ -974,7 +1047,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 == JaxCallableGraphMode.WARP_STAGED_EX and device.is_cuda: + elif self.graph_mode == JaxCallableGraphMode.WARP_STAGED_EX: # capturing with WARP using staging buffers and memcopies done outside of the graph wp_memcpy_batch = wp._src.context.runtime.core.wp_memcpy_batch @@ -1017,7 +1090,7 @@ class FfiCallable: # TODO: we should have a way of freeing this call_desc.capture = capture - elif self.graph_mode == JaxCallableGraphMode.WARP_STAGED and device.is_cuda: + elif self.graph_mode == JaxCallableGraphMode.WARP_STAGED: # 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 @@ -1224,6 +1297,7 @@ def jax_kernel( module_preload_mode=JaxModulePreloadMode.CURRENT_DEVICE, enable_backward: bool = False, has_side_effect: bool = False, + block_dim: int | None = None, ): """Create a JAX callback from a Warp kernel. @@ -1250,22 +1324,43 @@ def jax_kernel( kernel signature. The number of in-out arguments is included in ``num_outputs``. Not supported when ``enable_backward=True``. module_preload_mode: Specify the devices where the module should be preloaded. + ``JaxModulePreloadMode.ALL_DEVICES`` includes the host CPU and all + supported local CUDA devices. Preloading is best-effort: JAX devices + that cannot be mapped to Warp are skipped, and execution loads the + module for the callback's actual device. Module loading errors are + not suppressed. enable_backward: Enable automatic differentiation for this kernel. 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. + block_dim: Specify the number of threads per block for CUDA execution. + When ``None``, CUDA uses 256 threads per block. CPU execution always + uses one thread per block. The value is fixed when the wrapper is + constructed and is shared by forward and adjoint launches when + ``enable_backward=True``. Limitations: - All kernel arguments must be contiguous arrays or scalars. - Scalars must be static arguments in JAX. - 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. + - The CPU and CUDA backends are supported. JAX selects the backend from the + lowered computation's device. - ``output_dims`` and ``in_out_argnames`` are not supported when ``enable_backward=True``. """ check_jax_version() jax = _get_jax() + if block_dim is not None: + if isinstance(block_dim, bool): + raise TypeError("jax_kernel(): block_dim must be an integer or None.") + try: + block_dim = operator.index(block_dim) + except TypeError: + raise TypeError("jax_kernel(): block_dim must be an integer or None.") from None + if block_dim <= 0: + raise ValueError("jax_kernel(): block_dim must be positive.") + if isinstance(output_dims, dict): hashable_output_dims = tuple(sorted(output_dims.items())) elif hasattr(output_dims, "__len__"): @@ -1288,6 +1383,7 @@ def jax_kernel( hashable_output_dims, module_preload_mode, has_side_effect, + block_dim, ) with _FFI_REGISTRY_LOCK: @@ -1297,6 +1393,7 @@ def jax_kernel( num_outputs, vmap_method, launch_dims, + block_dim, output_dims, in_out_argnames, module_preload_mode, @@ -1342,6 +1439,7 @@ def jax_kernel( # 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 + _launch_block_dim = 256 if block_dim is None else block_dim def _resolve_launch_dims(call_args): if _user_launch_dims is not None: @@ -1362,7 +1460,13 @@ def jax_kernel( # Forward kernel wrapper: simply launches the kernel def fwd_kernel_wrapper(*args): - wp.launch(kernel, dim=_resolve_launch_dims(args), inputs=args[:num_inputs], outputs=args[num_inputs:]) + wp.launch( + kernel, + dim=_resolve_launch_dims(args), + inputs=args[:num_inputs], + outputs=args[num_inputs:], + block_dim=_launch_block_dim, + ) # update forward signature and annotations so jax_callable() sees a fully annotated function fwd_kernel_wrapper.__signature__ = signature @@ -1415,6 +1519,7 @@ def jax_kernel( adj_inputs=grad_in, adj_outputs=grad_out, adjoint=True, + block_dim=_launch_block_dim, ) # Build the backward wrapper signature expected by jax_callable @@ -1545,6 +1650,7 @@ def jax_kernel( # Reusing _user_launch_dims ensures int and 1-tuple forms of the same # value map to the same key. _user_launch_dims, + block_dim, ) if static_args: @@ -1585,6 +1691,7 @@ def jax_kernel( ) return cached(*args) + _checked_wrapper.block_dim = block_dim return _checked_wrapper @@ -1617,6 +1724,9 @@ def jax_callable( such as when the callable uses conditional graph nodes. ``JaxCallableGraphMode.NONE``: Disable graph capture. Use when the callable performs operations that are not legal in a graph, such as host synchronization. + On CPU, ``JaxCallableGraphMode.NONE`` and ``JaxCallableGraphMode.JAX`` + both execute without graph capture. The ``WARP``, ``WARP_STAGED``, and + ``WARP_STAGED_EX`` modes require CUDA. vmap_method: String specifying how the callback transforms under ``vmap()``. This argument can also be specified for individual calls. output_dims: Specify the default dimensions of output arrays. @@ -1632,6 +1742,11 @@ def jax_callable( 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. + ``JaxModulePreloadMode.ALL_DEVICES`` includes the host CPU and all + supported local CUDA devices. Preloading is best-effort: JAX devices + that cannot be mapped to Warp are skipped, and execution loads the + module for the callback's actual device. Module loading errors are + not suppressed. 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. @@ -1640,10 +1755,12 @@ def jax_callable( - Scalars must be static arguments in JAX. - Input and input-output arguments must precede the output arguments in the ``func`` definition. - There must be at least one output or input-output argument. - - Only the CUDA backend is supported. + - The CPU and CUDA backends are supported. JAX selects the backend from the + lowered computation's device. """ check_jax_version() + graph_mode = JaxCallableGraphMode(graph_mode) if isinstance(output_dims, dict): hashable_output_dims = tuple(sorted(output_dims.items())) @@ -1739,7 +1856,7 @@ def register_ffi_callback(name: str, func: Callable, graph_compatible: bool = Tr # TODO check that the name is not already registered - def ffi_callback(call_frame, platform="CUDA"): + def ffi_callback(call_frame, platform): try: extension = call_frame.contents.extension_start # On the first call, XLA runtime will query the API version and traits @@ -1751,7 +1868,7 @@ def register_ffi_callback(name: str, func: Callable, graph_compatible: bool = Tr metadata_ext = ctypes.cast(extension, ctypes.POINTER(XLA_FFI_Metadata_Extension)) metadata_ext.contents.metadata.contents.api_version.major_version = 0 metadata_ext.contents.metadata.contents.api_version.minor_version = 1 - if graph_compatible and platform == "CUDA": + if graph_compatible and platform == _FFI_PLATFORM_CUDA: # Turn on CUDA graphs for this handler. metadata_ext.contents.metadata.contents.traits = ( XLA_FFI_Handler_TraitsBits.COMMAND_BUFFER_COMPATIBLE @@ -1783,18 +1900,9 @@ def register_ffi_callback(name: str, func: Callable, graph_compatible: bool = Tr return None - FFI_CCALLFUNC = ctypes.CFUNCTYPE(ctypes.c_void_p, ctypes.POINTER(XLA_FFI_CallFrame)) - callback_func_cuda = FFI_CCALLFUNC(lambda call_frame: ffi_callback(call_frame, platform="CUDA")) - callback_func_host = FFI_CCALLFUNC(lambda call_frame: ffi_callback(call_frame, platform="Host")) + callback_funcs = _register_ffi_targets(name, ffi_callback) with _FFI_REGISTRY_LOCK: - _FFI_CALLBACK_REGISTRY[f"{name}_cuda"] = callback_func_cuda - _FFI_CALLBACK_REGISTRY[f"{name}_host"] = callback_func_host - ffi_ccall_address_cuda = ctypes.cast(callback_func_cuda, ctypes.c_void_p) - ffi_capsule_cuda = jax.ffi.pycapsule(ffi_ccall_address_cuda.value) - jax.ffi.register_ffi_target(name, ffi_capsule_cuda, platform="CUDA") - ffi_ccall_address_host = ctypes.cast(callback_func_host, ctypes.c_void_p) - ffi_capsule_host = jax.ffi.pycapsule(ffi_ccall_address_host.value) - jax.ffi.register_ffi_target(name, ffi_capsule_host, platform="Host") + _FFI_CALLBACK_REGISTRY[name] = callback_funcs ############################################################################### diff --git a/mjx/mujoco/mjx/third_party/warp/_src/jax/xla_ffi.py b/mjx/mujoco/mjx/third_party/warp/_src/jax/xla_ffi.py index a032c5b9..259e1c28 100644 --- a/mjx/mujoco/mjx/third_party/warp/_src/jax/xla_ffi.py +++ b/mjx/mujoco/mjx/third_party/warp/_src/jax/xla_ffi.py @@ -9,8 +9,6 @@ import numpy as np import warp as wp -_wp_module_name_ = "warp.jax.xla_ffi" - _xla_data_type_to_constructor = None _XLA_DATA_TYPE_TO_CONSTRUCTOR_LOCK = threading.Lock() diff --git a/mjx/pyproject.toml b/mjx/pyproject.toml index 86d502e4..2c07c4d6 100644 --- a/mjx/pyproject.toml +++ b/mjx/pyproject.toml @@ -37,7 +37,7 @@ dependencies = [ [project.optional-dependencies] warp = [ - "warp-lang==1.15.0", + "warp-lang==1.16.0", ] dev = [ "isort",