From 8eacacfd13a63407e901d8969f39ada7acaa3ac7 Mon Sep 17 00:00:00 2001 From: Baruch Tabanpour Date: Thu, 2 Oct 2025 16:07:03 -0700 Subject: [PATCH] Import NVIDIA/warp from GitHub. PiperOrigin-RevId: 814423562 Change-Id: Ibf441f161e25c51c494de84b0d263b32af76545c --- .../warp/jax_experimental/custom_call.py | 25 +++- .../third_party/warp/jax_experimental/ffi.py | 135 ++++++++++++++---- .../warp/jax_experimental/xla_ffi.py | 23 ++- 3 files changed, 150 insertions(+), 33 deletions(-) diff --git a/mjx/mujoco/mjx/third_party/warp/jax_experimental/custom_call.py b/mjx/mujoco/mjx/third_party/warp/jax_experimental/custom_call.py index a90f09a2..ff8cc0c7 100644 --- a/mjx/mujoco/mjx/third_party/warp/jax_experimental/custom_call.py +++ b/mjx/mujoco/mjx/third_party/warp/jax_experimental/custom_call.py @@ -19,6 +19,7 @@ import warp as wp from warp.context import type_str from warp.jax import get_jax_device from warp.types import array_t, launch_bounds_t, strides_from_shape +from warp.utils import warn _jax_warp_p = None @@ -28,7 +29,7 @@ _registered_kernels = [None] _registered_kernel_to_id = {} -def jax_kernel(kernel, launch_dims=None): +def jax_kernel(kernel, launch_dims=None, quiet=False): """Create a Jax primitive from a Warp kernel. NOTE: This is an experimental feature under development. @@ -38,6 +39,7 @@ def jax_kernel(kernel, launch_dims=None): launch_dims: Optional. Specify the kernel launch dimensions. If None, dimensions are inferred from the shape of the first argument. This option when set will specify the output dimensions. + quiet: Optional. If True, suppress deprecation warnings with newer JAX versions. Limitations: - All kernel arguments must be contiguous arrays. @@ -46,6 +48,27 @@ def jax_kernel(kernel, launch_dims=None): - Only the CUDA backend is supported. """ + import jax + + # check if JAX version supports this + if jax.__version_info__ < (0, 4, 25) or jax.__version_info__ >= (0, 8, 0): + msg = ( + "This version of jax_kernel() requires JAX version 0.4.25 - 0.7.x, " + 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." + raise RuntimeError(msg) + + # deprecation warning + if jax.__version_info__ >= (0, 5, 0) and not quiet: + warn( + "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). " + "In Warp release 1.10, the FFI version will become the default implementation of jax_kernel().", + DeprecationWarning, + ) + if _jax_warp_p is None: # Create and register the primitive _create_jax_warp_primitive() diff --git a/mjx/mujoco/mjx/third_party/warp/jax_experimental/ffi.py b/mjx/mujoco/mjx/third_party/warp/jax_experimental/ffi.py index 05165b49..b31d1a7b 100644 --- a/mjx/mujoco/mjx/third_party/warp/jax_experimental/ffi.py +++ b/mjx/mujoco/mjx/third_party/warp/jax_experimental/ffi.py @@ -13,6 +13,7 @@ # See the License for the specific language governing permissions and # limitations under the License. +import collections import ctypes import threading import traceback @@ -28,6 +29,24 @@ from warp.types import array_t, launch_bounds_t, strides_from_shape, type_to_war from .xla_ffi import * +jax_callable_default_graph_cache_max: int | None = 32 +""" +Maximum size of the graph cache for graphs captured using ``GraphMode.WARP``, unlimited if ``None``. +Example usage: ``warp.jax_experimental.ffi.jax_callable_default_graph_cache_max = 42``. +""" + + +def check_jax_version(): + # check if JAX version supports this + if jax.__version_info__ < (0, 5, 0): + msg = ( + "This version of jax_kernel() requires JAX version 0.5.0 or higher, " + f"but installed JAX version is {jax.__version_info__}." + ) + if jax.__version_info__ >= (0, 4, 25): + msg += " Please use warp.jax_experimental.custom_call.jax_kernel instead." + raise RuntimeError(msg) + class GraphMode(IntEnum): NONE = 0 # don't capture a graph @@ -338,11 +357,10 @@ class FfiKernel: class FfiCallDesc: def __init__(self, static_inputs): self.static_inputs = static_inputs - self.captures = {} class FfiCallable: - def __init__(self, func, num_outputs, graph_mode, vmap_method, output_dims, in_out_argnames): + def __init__(self, func, num_outputs, graph_mode, vmap_method, output_dims, in_out_argnames, graph_cache_max): self.func = func self.name = generate_unique_name(func) self.num_outputs = num_outputs @@ -353,6 +371,10 @@ class FfiCallable: self.call_id = 0 self.call_descriptors = {} + # LRU cache of graphs captured by Warp + self._graph_cache_max = graph_cache_max + self.captures = collections.OrderedDict() + in_out_argnames_list = in_out_argnames or [] in_out_argnames = set(in_out_argnames_list) if len(in_out_argnames_list) != len(in_out_argnames): @@ -559,8 +581,8 @@ class FfiCallable: # 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] - buffer_hash = hash((*ip, *op)) - capture = call_desc.captures.get(buffer_hash) + capture_key = hash((call_id, *ip, *op)) + capture = self.captures.get(capture_key) # launch existing graph if capture is not None: @@ -578,6 +600,9 @@ class FfiCallable: if not wp.context.runtime.core.wp_cuda_graph_launch(graph.graph_exec, cuda_stream): raise RuntimeError(f"Graph launch error: {wp.context.runtime.get_error_string()}") + # update the graph cache to keep recently used graphs alive + self.captures.move_to_end(capture_key) + # early out return @@ -620,7 +645,10 @@ class FfiCallable: self.func(*arg_list) wp.capture_launch(capture.graph) # keep a reference to the capture object and reuse it with same buffers - call_desc.captures[buffer_hash] = capture + self.captures[capture_key] = capture + # respect the cache size limit if specified + if self._graph_cache_max is not None and len(self.captures) > self._graph_cache_max: + self.captures.popitem(last=False) else: # not capturing self.func(*arg_list) @@ -633,10 +661,28 @@ class FfiCallable: return None + @property + def graph_cache_max(self) -> int | None: + return self._graph_cache_max + + @graph_cache_max.setter + def graph_cache_max(self, value: int | None): + if value != self._graph_cache_max: + if value is not None and (self._graph_cache_max is None or value < self._graph_cache_max): + # trim the cache if needed + while len(self.captures) > value: + self.captures.popitem(last=False) + self._graph_cache_max = value + + @property + def graph_cache_size(self) -> int: + return len(self.captures) + # Holders for the custom callbacks to keep them alive. -_FFI_CALLABLE_REGISTRY: dict[str, FfiCallable] = {} _FFI_KERNEL_REGISTRY: dict[str, FfiKernel] = {} +_FFI_CALLABLE_REGISTRY: dict[str, FfiCallable] = {} +_FFI_CALLBACK_REGISTRY: dict[str, ctypes.CFUNCTYPE] = {} _FFI_REGISTRY_LOCK = threading.Lock() @@ -649,17 +695,17 @@ def jax_kernel( Args: kernel: The Warp kernel to launch. - num_outputs: Optional. Specify the number of output arguments if greater than 1. + num_outputs: Specify the number of output arguments if greater than 1. This must include the number of ``in_out_arguments``. - vmap_method: Optional. String specifying how the callback transforms under ``vmap()``. + vmap_method: String specifying how the callback transforms under ``vmap()``. This argument can also be specified for individual calls. - launch_dims: Optional. Specify the default kernel launch dimensions. If None, launch + 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. - output_dims: Optional. Specify the default dimensions of output arrays. If None, output + 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. - in_out_argnames: Optional. Names of input-output arguments. + in_out_argnames: Names of input-output arguments. Limitations: - All kernel arguments must be contiguous arrays or scalars. @@ -668,8 +714,12 @@ def jax_kernel( - There must be at least one output or input-output argument. - Only the CUDA backend is supported. """ + + check_jax_version() + key = ( kernel.func, + kernel.sig, num_outputs, vmap_method, tuple(launch_dims) if launch_dims else launch_dims, @@ -692,6 +742,7 @@ def jax_callable( vmap_method: Optional[str] = "broadcast_all", output_dims=None, in_out_argnames=None, + graph_cache_max: int | None = None, ): """Create a JAX callback from an annotated Python function. @@ -701,22 +752,24 @@ def jax_callable( Args: func: The Python function to call. - num_outputs: Optional. Specify the number of output arguments if greater than 1. + num_outputs: Specify the number of output arguments if greater than 1. This must include the number of ``in_out_arguments``. - graph_compatible: Optional. Whether the function can be called during CUDA graph capture. + graph_compatible: Whether the function can be called during CUDA graph capture. This argument is deprecated, use ``graph_mode`` instead. - graph_mode: Optional. CUDA graph capture mode. - ``GraphMode.JAX`` (default): Let JAX capture the graph, which may be used as a subgraph in an enclosing capture. - ``GraphMode.WARP``: Let Warp capture the graph. Use this mode when the callable cannot be used as a subraph, + 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, 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, such as host synchronization. - vmap_method: Optional. String specifying how the callback transforms under ``vmap()``. + vmap_method: String specifying how the callback transforms under ``vmap()``. This argument can also be specified for individual calls. - output_dims: Optional. Specify the default dimensions of output arrays. + 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. - in_out_argnames: Optional. Names of input-output arguments. + in_out_argnames: Names of input-output arguments. + graph_cache_max: Maximum number of cached graphs captured using ``GraphMode.WARP``. + If ``None``, use ``warp.jax_experimental.ffi.jax_callable_default_graph_cache_max``. Limitations: - All kernel arguments must be contiguous arrays or scalars. @@ -726,6 +779,8 @@ def jax_callable( - Only the CUDA backend is supported. """ + check_jax_version() + if graph_compatible is not None: wp.utils.warn( "The `graph_compatible` argument is deprecated, use `graph_mode` instead.", @@ -735,6 +790,10 @@ def jax_callable( if graph_compatible is False: graph_mode = GraphMode.NONE + if graph_cache_max is None: + graph_cache_max = jax_callable_default_graph_cache_max + + # Note: we don't include graph_cache_max in the key, it is applied below. key = ( func, num_outputs, @@ -744,11 +803,35 @@ def jax_callable( ) with _FFI_REGISTRY_LOCK: - if key not in _FFI_CALLABLE_REGISTRY: - new_callable = FfiCallable(func, num_outputs, graph_mode, vmap_method, output_dims, in_out_argnames) - _FFI_CALLABLE_REGISTRY[key] = new_callable + callable = _FFI_CALLABLE_REGISTRY.get(key) + if callable is None: + callable = FfiCallable( + func, + num_outputs, + graph_mode, + vmap_method, + output_dims, + in_out_argnames, + graph_cache_max, + ) + _FFI_CALLABLE_REGISTRY[key] = callable + else: + # make sure we're using the latest graph cache max + callable.graph_cache_max = graph_cache_max - return _FFI_CALLABLE_REGISTRY[key] + return callable + + +def clear_jax_callable_graph_cache(callable: FfiCallable | None = None): + """Clear the graph cache of the given callable or all callables if ``None``.""" + + if callable is not None: + callable.captures.clear() + else: + # apply to all callables + with _FFI_REGISTRY_LOCK: + for callable in _FFI_CALLABLE_REGISTRY.values(): + callable.captures.clear() ############################################################################### @@ -769,9 +852,11 @@ def register_ffi_callback(name: str, func: Callable, graph_compatible: bool = Tr Args: name: A unique FFI callback name. func: The Python function to call. - graph_compatible: Optional. Whether the function can be called during CUDA graph capture. + graph_compatible: Whether the function can be called during CUDA graph capture. """ + check_jax_version() + # TODO check that the name is not already registered def ffi_callback(call_frame): @@ -817,7 +902,7 @@ def register_ffi_callback(name: str, func: Callable, graph_compatible: bool = Tr FFI_CCALLFUNC = ctypes.CFUNCTYPE(ctypes.c_void_p, ctypes.POINTER(XLA_FFI_CallFrame)) callback_func = FFI_CCALLFUNC(ffi_callback) with _FFI_REGISTRY_LOCK: - _FFI_CALLABLE_REGISTRY[name] = callback_func + _FFI_CALLBACK_REGISTRY[name] = callback_func ffi_ccall_address = ctypes.cast(callback_func, ctypes.c_void_p) ffi_capsule = jax.ffi.pycapsule(ffi_ccall_address.value) jax.ffi.register_ffi_target(name, ffi_capsule, platform="CUDA") diff --git a/mjx/mujoco/mjx/third_party/warp/jax_experimental/xla_ffi.py b/mjx/mujoco/mjx/third_party/warp/jax_experimental/xla_ffi.py index 8f33030b..8e311143 100644 --- a/mjx/mujoco/mjx/third_party/warp/jax_experimental/xla_ffi.py +++ b/mjx/mujoco/mjx/third_party/warp/jax_experimental/xla_ffi.py @@ -475,17 +475,26 @@ _xla_data_type_to_constructor = { XLA_FFI_DataType.C64: jnp.complex64, XLA_FFI_DataType.C128: jnp.complex128, # XLA_FFI_DataType.TOKEN - XLA_FFI_DataType.F8E5M2: jnp.float8_e5m2, - XLA_FFI_DataType.F8E3M4: jnp.float8_e3m4, - XLA_FFI_DataType.F8E4M3: jnp.float8_e4m3, - XLA_FFI_DataType.F8E4M3FN: jnp.float8_e4m3fn, - XLA_FFI_DataType.F8E4M3B11FNUZ: jnp.float8_e4m3b11fnuz, - XLA_FFI_DataType.F8E5M2FNUZ: jnp.float8_e5m2fnuz, - XLA_FFI_DataType.F8E4M3FNUZ: jnp.float8_e4m3fnuz, # 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