diff --git a/mjx/cuda_requirements.txt b/mjx/cuda_requirements.txt index b2e81f22..deca223c 100644 --- a/mjx/cuda_requirements.txt +++ b/mjx/cuda_requirements.txt @@ -16,12 +16,8 @@ jax-cuda12-pjrt==0.5.3; python_version >= '3.10' \ jax-cuda12-pjrt==0.4.30; python_version == '3.9' \ --hash=sha256:895d0198ad99638fcaf976c47592e2a543eef79ea15fabd24a402d055390c328 \ --hash=sha256:c36fb1e0c236563bf3a87e70f4d1ab28a31d7cf5d722c9ede30c4172116e8bcb -warp-lang==1.10.1 \ - --hash=sha256:0c6f44d4136cfc86316f5d35883a863b4dca2bc318331597e755ea872db0cd97 \ - --hash=sha256:2884b642f16b07b930b3605193c1b97d183eef80c5b0083d3a17473aee92138e \ - --hash=sha256:7c8f839a01042677d1f417d31abe9304f9319fca7041c2ed82937ed511eb2362 \ - --hash=sha256:7addd14a913b50c406ba3b1abebf3a208d01591a1c382da799cff2806d234c5a \ - --hash=sha256:5f5ce0147d48f86ccab66b834aa54c8adedbbb2ba86e45b2b904b182fa963509 \ - --hash=sha256:2068e877dd2972d841c1eb63f98eadce326e77ba58f09e8474d0847276f47633 \ - --hash=sha256:96d24694333e30eac888b2c2962bc5ca4b2369f56707e4c54a465fa37b7ae007 \ - --hash=sha256:b77ffe935e06cb57e76f6d4dd8bae7d59294ce63977ea8a0a4fd8e30d77c0a7f +warp-lang==1.11.0 \ + --hash=sha256:3a4f1c9a6e721d7de7d6dad6b242c54afaf20c6e14a767c0da03e5e963fcc13c \ + --hash=sha256:524dce20de6162ba25333552168ebf430973050e00d9f8116b8df41a60d25d6e \ + --hash=sha256:1ae6cfc226107f96e4d495b41a3dab32488e8ee8f074b0e1bcaf22e7fb8c904d \ + --hash=sha256:80d8493cbe243a3510134f3af289646d7bd7484217a30ecf565d676466ef8a5e diff --git a/mjx/mujoco/mjx/third_party/warp/_src/jax_experimental/custom_call.py b/mjx/mujoco/mjx/third_party/warp/_src/jax_experimental/custom_call.py index cd376477..b46a1071 100644 --- a/mjx/mujoco/mjx/third_party/warp/_src/jax_experimental/custom_call.py +++ b/mjx/mujoco/mjx/third_party/warp/_src/jax_experimental/custom_call.py @@ -14,6 +14,7 @@ # limitations under the License. import ctypes +from functools import reduce import warp as wp from warp._src.context import type_str @@ -34,14 +35,23 @@ _registered_kernel_to_id = {} 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. + .. 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 + 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`. Args: kernel: The Warp kernel to be wrapped. - launch_dims: Optional. Specify the kernel launch dimensions. If None, + launch_dims: 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. + quiet: If ``True``, suppress deprecation warnings with newer JAX versions. + + Raises: + RuntimeError: If JAX version is < 0.4.25 or >= 0.8.0. Limitations: - All kernel arguments must be contiguous arrays. @@ -49,8 +59,7 @@ def jax_kernel(kernel, launch_dims=None, quiet=False): - There must be at least one input argument and at least one output argument. - Only the CUDA backend is supported. """ - - import jax + import jax # noqa: PLC0415 # check if JAX version supports this if jax.__version_info__ < (0, 4, 25) or jax.__version_info__ >= (0, 8, 0): @@ -139,13 +148,11 @@ def _warp_custom_callback(stream, buffers, opaque, opaque_len): def _create_jax_warp_primitive(): - from functools import reduce - - import jax - from jax._src.interpreters import batching - from jax.interpreters import mlir - from jax.interpreters.mlir import ir - from jaxlib.hlo_helpers import custom_call + import jax # noqa: PLC0415 + from jax._src.interpreters import batching # noqa: PLC0415 + from jax.interpreters import mlir # noqa: PLC0415 + from jax.interpreters.mlir import ir # noqa: PLC0415 + from jaxlib.hlo_helpers import custom_call # noqa: PLC0415 global _jax_warp_p global _cc_callback @@ -154,7 +161,7 @@ def _create_jax_warp_primitive(): # TODO add default implementation that calls the kernel via warp. try: # newer JAX versions - import jax.extend + import jax.extend # noqa: PLC0415 _jax_warp_p = jax.extend.core.Primitive("jax_warp") except (ImportError, AttributeError): @@ -201,7 +208,7 @@ def _create_jax_warp_primitive(): raise Exception(f"Argument {warp_arg.label} has too few non-matrix/vector dimensions") index_rest = len(actual_shape) - warp_arg.type.ndim + 1 leading_size = reduce(lambda x, y: x * y, actual_shape[:index_rest]) - return [leading_size] + actual_shape[index_rest:] + return [leading_size, *actual_shape[index_rest:]] # Infer array dimensions from input type. def infer_dimensions(warp_arg, actual_shape): diff --git a/mjx/mujoco/mjx/third_party/warp/_src/jax_experimental/ffi.py b/mjx/mujoco/mjx/third_party/warp/_src/jax_experimental/ffi.py index 64592ac0..43f63397 100644 --- a/mjx/mujoco/mjx/third_party/warp/_src/jax_experimental/ffi.py +++ b/mjx/mujoco/mjx/third_party/warp/_src/jax_experimental/ffi.py @@ -13,20 +13,23 @@ # See the License for the specific language governing permissions and # limitations under the License. +from __future__ import annotations + import collections import ctypes import inspect import threading import traceback from enum import IntEnum -from typing import Callable, Optional +from typing import Callable 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.types import array_t, launch_bounds_t, strides_from_shape, type_to_warp +from warp._src.types import array_t, launch_bounds_t, strides_from_shape, type_size_in_bytes, type_to_warp from .xla_ffi import * @@ -36,9 +39,9 @@ _wp_module_name_ = "warp.jax_experimental.ffi" DiffKernelCacheKey = tuple[Callable, tuple, int, str, tuple[str, ...]] # Holders for the custom callbacks to keep them alive. -_FFI_KERNEL_REGISTRY: dict[str, "FfiKernel"] = {} +_FFI_KERNEL_REGISTRY: dict[str, FfiKernel] = {} _FFI_DIFF_KERNEL_REGISTRY: dict[DiffKernelCacheKey, Callable] = {} -_FFI_CALLABLE_REGISTRY: dict[str, "FfiCallable"] = {} +_FFI_CALLABLE_REGISTRY: dict[str, FfiCallable] = {} _FFI_CALLBACK_REGISTRY: dict[str, ctypes.CFUNCTYPE] = {} _FFI_REGISTRY_LOCK = threading.Lock() @@ -62,6 +65,8 @@ class GraphMode(IntEnum): NONE = 0 # don't capture a graph JAX = 1 # let JAX capture a graph WARP = 2 # let Warp capture a graph + WARP_STAGED = 3 # use Warp graph with staging buffers, copy inside of the graph + WARP_STAGED_EX = 4 # use Warp graph with staging buffers, copy outside of the graph class ModulePreloadMode(IntEnum): @@ -390,6 +395,28 @@ class FfiKernel: class FfiCallDesc: def __init__(self, static_inputs): self.static_inputs = static_inputs + self.capture = None + + # staging arrays + self.input_staging_arrays = None # inputs copied on each call + self.output_staging_arrays = None # outputs copied on each call + self.static_staging_arrays = None # arrays copied only once + + # input memcpy info + self.input_memcpy_count = 0 + self.input_memcpy_indices = None # indices in FFI input buffers + self.input_memcpy_srcs = None + self.input_memcpy_dsts = None + self.input_memcpy_sizes = None + self.input_memcpy_kinds = None + + # output memcpy info + self.output_memcpy_count = 0 + self.output_memcpy_indices = None # indices in FFI output buffers + self.output_memcpy_srcs = None + self.output_memcpy_dsts = None + self.output_memcpy_sizes = None + self.output_memcpy_kinds = None class FfiCallable: @@ -403,6 +430,8 @@ class FfiCallable: vmap_method, output_dims, in_out_argnames, + stage_in_argnames, + stage_out_argnames, graph_cache_max, module_preload_mode, ): @@ -421,6 +450,10 @@ class FfiCallable: self._graph_cache_max = graph_cache_max self.captures = collections.OrderedDict() + # Selective staging: None means copy all, else copy only specified + self.stage_in_argnames = set(stage_in_argnames) if stage_in_argnames else None + self.stage_out_argnames = set(stage_out_argnames) if stage_out_argnames else None + 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): @@ -444,7 +477,10 @@ class FfiCallable: # parse type annotations self.args = [] + self.arg_input_indices = [None] * num_args # index in FFI input buffers + self.arg_output_indices = [None] * num_args # index in FFI output buffers arg_idx = 0 + output_idx = 0 for arg_name, arg_type in argspec.annotations.items(): if arg_name == "return": if arg_type is not None: @@ -452,7 +488,7 @@ class FfiCallable: continue else: arg = FfiArg(arg_name, arg_type, arg_name in in_out_argnames) - if arg_name in in_out_argnames: + if arg.in_out: in_out_argnames.remove(arg_name) if arg.is_array: if arg_idx < self.num_inputs and self.first_array_arg is None: @@ -465,6 +501,14 @@ class FfiCallable: " in_out arguments should be placed before output-only arguments." ) + # map each argument index to FFI input/output buffer indices + # (including in-out arguments) + if arg_idx < self.num_inputs: + self.arg_input_indices[arg_idx] = arg_idx + if arg_idx >= self.num_inputs or arg.in_out: + self.arg_output_indices[arg_idx] = output_idx + output_idx += 1 + arg_idx += 1 if in_out_argnames: @@ -646,19 +690,9 @@ class FfiCallable: # launch existing graph if capture is not None: - # NOTE: We use the native graph API to avoid overhead with obtaining Stream and Device objects in Python. - # This code should match wp.capture_launch(). - graph = capture.graph - if graph.graph_exec is None: - g = ctypes.c_void_p() - if not wp._src.context.runtime.core.wp_cuda_graph_create_exec( - graph.device.context, cuda_stream, graph.graph, ctypes.byref(g) - ): - raise RuntimeError(f"Graph creation error: {wp.context.runtime.get_error_string()}") - graph.graph_exec = g - - if not wp._src.context.runtime.core.wp_cuda_graph_launch(graph.graph_exec, cuda_stream): - raise RuntimeError(f"Graph launch error: {wp.context.runtime.get_error_string()}") + graph_exec = capture.graph.graph_exec + if not wp._src.context.runtime.core.wp_cuda_graph_launch(graph_exec, cuda_stream): + raise RuntimeError(f"Graph launch error: {wp._src.context.runtime.get_error_string()}") # update the graph cache to keep recently used graphs alive self.captures.move_to_end(capture_key) @@ -666,6 +700,89 @@ class FfiCallable: # early out return + elif self.graph_mode == GraphMode.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 + wp_memcpy_batch = wp._src.context.runtime.core.wp_memcpy_batch + + # set source pointers for input memcopies + for memcpy_idx, input_idx in enumerate(call_desc.input_memcpy_indices): + call_desc.input_memcpy_srcs[memcpy_idx] = inputs[input_idx].contents.data + + # copy inputs to staging buffers + if not wp_memcpy_batch( + context, + call_desc.input_memcpy_dsts, + call_desc.input_memcpy_srcs, + call_desc.input_memcpy_sizes, + call_desc.input_memcpy_count, + cuda_stream, + ): + raise RuntimeError( + f"Failed to run input memcpy batch: {wp._src.context.runtime.get_error_string()}" + ) + + # launch existing graph + if not wp._src.context.runtime.core.wp_cuda_graph_launch(graph_exec, cuda_stream): + raise RuntimeError(f"Graph launch error: {wp._src.context.runtime.get_error_string()}") + + # set destination pointers for output memcopies + for memcpy_idx, output_idx in enumerate(call_desc.output_memcpy_indices): + call_desc.output_memcpy_dsts[memcpy_idx] = outputs[output_idx].contents.data + + # copy the outputs from staging buffers + if not wp_memcpy_batch( + context, + call_desc.output_memcpy_dsts, + call_desc.output_memcpy_srcs, + call_desc.output_memcpy_sizes, + call_desc.output_memcpy_count, + cuda_stream, + ): + raise RuntimeError( + f"Failed to run output memcpy batch: {wp._src.context.runtime.get_error_string()}" + ) + + # early out + return + + elif self.graph_mode == GraphMode.WARP_STAGED: + if call_desc.capture is not None: + graph_exec = call_desc.capture.graph.graph_exec + + # set source pointers for input memcpy nodes + memcpy_idx = 0 + for input_idx in call_desc.input_memcpy_indices: + call_desc.memcpy_srcs[memcpy_idx] = inputs[input_idx].contents.data + memcpy_idx += 1 + + # set destination pointers for output memcpy nodes + for output_idx in call_desc.output_memcpy_indices: + call_desc.memcpy_dsts[memcpy_idx] = outputs[output_idx].contents.data + memcpy_idx += 1 + + # update all memcpy nodes + if not wp._src.context.runtime.core.wp_cuda_graph_update_memcpy_batch( + graph_exec, + call_desc.memcpy_nodes, + call_desc.memcpy_dsts, + call_desc.memcpy_srcs, + call_desc.memcpy_sizes, + call_desc.memcpy_kinds, + len(call_desc.memcpy_nodes), + ): + raise RuntimeError( + f"Failed to update graph memcpy batch: {wp._src.context.runtime.get_error_string()}" + ) + + # launch existing graph + if not wp._src.context.runtime.core.wp_cuda_graph_launch(graph_exec, cuda_stream): + raise RuntimeError(f"Graph launch error: {wp._src.context.runtime.get_error_string()}") + + # early out + return + device_ordinal = get_device_ordinal_from_callframe(call_frame.contents) device = wp.get_cuda_device(device_ordinal) stream = wp.Stream(device, cuda_stream=cuda_stream) @@ -698,18 +815,143 @@ class FfiCallable: # capturing with JAX with wp.ScopedCapture(external=True) as capture: self.func(*arg_list) + # keep a reference to the capture object to prevent required modules getting unloaded call_desc.capture = capture + elif self.graph_mode == GraphMode.WARP: # capturing with WARP with wp.ScopedCapture() as capture: self.func(*arg_list) wp.capture_launch(capture.graph) + # keep a reference to the capture object and reuse it with same buffers 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) + + elif self.graph_mode == GraphMode.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 + + # prepare staging arrays and memcpy info + self._prepare_staging(arg_list, call_desc) + + # copy inputs to staging arrays (including in-out arrays) + if not wp_memcpy_batch( + device.context, + call_desc.input_memcpy_dsts, + call_desc.input_memcpy_srcs, + call_desc.input_memcpy_sizes, + call_desc.input_memcpy_count, + cuda_stream, + ): + raise RuntimeError( + f"Failed to run input memcpy batch: {wp._src.context.runtime.get_error_string()}" + ) + + # capture callback using staging arrays + with wp.ScopedCapture() as capture: + self.func(*arg_list) + + wp.capture_launch(capture.graph) + + # copy outputs from staging arrays (including in-out arrays) + if not wp_memcpy_batch( + device.context, + call_desc.output_memcpy_dsts, + call_desc.output_memcpy_srcs, + call_desc.output_memcpy_sizes, + call_desc.output_memcpy_count, + cuda_stream, + ): + raise RuntimeError( + f"Failed to run output memcpy batch: {wp._src.context.runtime.get_error_string()}" + ) + + # save the capture for replays + # TODO: we should have a way of freeing this + call_desc.capture = capture + + elif self.graph_mode == GraphMode.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 + ) + + # prepare staging arrays and memcpy info + self._prepare_staging(arg_list, call_desc) + + # prepare graph memcpy nodes + input_memcpy_count = call_desc.input_memcpy_count + output_memcpy_count = call_desc.output_memcpy_count + input_memcpy_nodes = (ctypes.c_void_p * input_memcpy_count)() + output_memcpy_nodes = (ctypes.c_void_p * output_memcpy_count)() + + # capture using staging arrays and include memory copies + with wp.ScopedCapture() as capture: + # copy inputs + if not wp_cuda_graph_insert_memcpy_batch( + device.context, + cuda_stream, + call_desc.input_memcpy_dsts, + call_desc.input_memcpy_srcs, + call_desc.input_memcpy_sizes, + call_desc.input_memcpy_kinds, + call_desc.input_memcpy_count, + input_memcpy_nodes, + ): + raise RuntimeError( + f"Failed to insert input memcpy batch: {wp._src.context.runtime.get_error_string()}" + ) + + # run the callback + self.func(*arg_list) + + # copy outputs + if not wp_cuda_graph_insert_memcpy_batch( + device.context, + cuda_stream, + call_desc.output_memcpy_dsts, + call_desc.output_memcpy_srcs, + call_desc.output_memcpy_sizes, + call_desc.output_memcpy_kinds, + call_desc.output_memcpy_count, + output_memcpy_nodes, + ): + raise RuntimeError( + f"Failed to insert output memcpy batch: {wp._src.context.runtime.get_error_string()}" + ) + + wp.capture_launch(capture.graph) + + # concatenate input and output memcopy nodes so they can be updated in one call + num_nodes = input_memcpy_count + output_memcpy_count + call_desc.memcpy_nodes = (ctypes.c_void_p * num_nodes)() + call_desc.memcpy_srcs = (ctypes.c_void_p * num_nodes)() + call_desc.memcpy_dsts = (ctypes.c_void_p * num_nodes)() + call_desc.memcpy_sizes = (ctypes.c_size_t * num_nodes)() + call_desc.memcpy_kinds = (ctypes.c_int * num_nodes)() + for i in range(input_memcpy_count): + call_desc.memcpy_nodes[i] = input_memcpy_nodes[i] + call_desc.memcpy_srcs[i] = call_desc.input_memcpy_srcs[i] + call_desc.memcpy_dsts[i] = call_desc.input_memcpy_dsts[i] + call_desc.memcpy_sizes[i] = call_desc.input_memcpy_sizes[i] + call_desc.memcpy_kinds[i] = call_desc.input_memcpy_kinds[i] + for i in range(output_memcpy_count): + j = input_memcpy_count + i + call_desc.memcpy_nodes[j] = output_memcpy_nodes[i] + call_desc.memcpy_srcs[j] = call_desc.output_memcpy_srcs[i] + call_desc.memcpy_dsts[j] = call_desc.output_memcpy_dsts[i] + call_desc.memcpy_sizes[j] = call_desc.output_memcpy_sizes[i] + call_desc.memcpy_kinds[j] = call_desc.output_memcpy_kinds[i] + + # save the capture for replays + # TODO: we should have a way of freeing this + call_desc.capture = capture + else: # not capturing self.func(*arg_list) @@ -722,6 +964,96 @@ class FfiCallable: return None + def _prepare_staging(self, arg_list, call_desc): + # create staging arrays + input_callback_arrays = [] + input_staging_arrays = [] + input_memcpy_indices = [] + output_callback_arrays = [] + output_staging_arrays = [] + output_memcpy_indices = [] + static_staging_arrays = [] + for i, arg in enumerate(arg_list): + # we only care about arrays with non-zero size + if isinstance(arg, wp.array) and arg.size > 0: + staging_arr = wp.empty_like(arg) + input_idx = self.arg_input_indices[i] + if input_idx is not None: + # check if this input needs to be copied every time or just once + if self.stage_in_argnames is None or self.args[i].name in self.stage_in_argnames: + input_callback_arrays.append(arg) + input_staging_arrays.append(staging_arr) + input_memcpy_indices.append(input_idx) + else: + wp.copy(staging_arr, arg) + static_staging_arrays.append(staging_arr) + output_idx = self.arg_output_indices[i] + if output_idx is not None: + # check if this output needs to be copied every time or just once + if self.stage_out_argnames is None or self.args[i].name in self.stage_out_argnames: + output_callback_arrays.append(arg) + output_staging_arrays.append(staging_arr) + output_memcpy_indices.append(output_idx) + else: + wp.copy(staging_arr, arg) + static_staging_arrays.append(staging_arr) + # substitute staging array in argument list + arg_list[i] = staging_arr + + # prepare input memcpy batch + input_memcpy_count = len(input_staging_arrays) + call_desc.input_memcpy_count = input_memcpy_count + call_desc.input_memcpy_indices = input_memcpy_indices + call_desc.input_memcpy_srcs = (ctypes.c_void_p * input_memcpy_count)() + call_desc.input_memcpy_dsts = (ctypes.c_void_p * input_memcpy_count)() + call_desc.input_memcpy_sizes = (ctypes.c_size_t * input_memcpy_count)() + call_desc.input_memcpy_kinds = (ctypes.c_int * input_memcpy_count)() + for i in range(input_memcpy_count): + size = input_staging_arrays[i].size * type_size_in_bytes(input_staging_arrays[i].dtype) + call_desc.input_memcpy_srcs[i] = input_callback_arrays[i].ptr + call_desc.input_memcpy_dsts[i] = input_staging_arrays[i].ptr + call_desc.input_memcpy_sizes[i] = size + call_desc.input_memcpy_kinds[i] = CudaMemcpyKind.D2D + + # prepare output memcpy batch + output_memcpy_count = len(output_staging_arrays) + call_desc.output_memcpy_count = output_memcpy_count + call_desc.output_memcpy_indices = output_memcpy_indices + call_desc.output_memcpy_srcs = (ctypes.c_void_p * output_memcpy_count)() + call_desc.output_memcpy_dsts = (ctypes.c_void_p * output_memcpy_count)() + call_desc.output_memcpy_sizes = (ctypes.c_size_t * output_memcpy_count)() + call_desc.output_memcpy_kinds = (ctypes.c_int * output_memcpy_count)() + for i in range(output_memcpy_count): + size = output_staging_arrays[i].size * type_size_in_bytes(output_staging_arrays[i].dtype) + call_desc.output_memcpy_srcs[i] = output_staging_arrays[i].ptr + call_desc.output_memcpy_dsts[i] = output_callback_arrays[i].ptr + call_desc.output_memcpy_sizes[i] = size + call_desc.output_memcpy_kinds[i] = CudaMemcpyKind.D2D + + # hang on to the staging arrays to prevent GC + # TODO: we should have a way of freeing this + call_desc.input_staging_arrays = input_staging_arrays + call_desc.output_staging_arrays = output_staging_arrays + call_desc.static_staging_arrays = static_staging_arrays + + if wp.config.verbose: + # print some stats + total_input_size = 0 + for i in range(input_memcpy_count): + total_input_size += int(call_desc.input_memcpy_sizes[i]) + total_output_size = 0 + for i in range(output_memcpy_count): + total_output_size += int(call_desc.output_memcpy_sizes[i]) + print("FFI graph staging stats:") + print(f" input memcpy indices: {input_memcpy_indices}") + print(f" output memcpy indices: {output_memcpy_indices}") + print(f" input memcpy count: {input_memcpy_count}") + print(f" output memcpy count: {output_memcpy_count}") + print(f" total memcpy count: {input_memcpy_count + output_memcpy_count}") + print(f" total input size: {total_input_size} bytes") + print(f" total output size: {total_output_size} bytes") + print(f" total size: {total_input_size + total_output_size} bytes") + @property def graph_cache_max(self) -> int | None: return self._graph_cache_max @@ -1039,11 +1371,12 @@ def jax_kernel( def jax_callable( func: Callable, num_outputs: int = 1, - graph_compatible: Optional[bool] = None, # deprecated graph_mode: GraphMode = GraphMode.JAX, - vmap_method: Optional[str] = "broadcast_all", + 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, ): @@ -1057,8 +1390,6 @@ def jax_callable( func: The Python function to call. num_outputs: Specify the number of output arguments if greater than 1. This must include the number of ``in_out_arguments``. - graph_compatible: Whether the function can be called during CUDA graph capture. - This argument is deprecated, use ``graph_mode`` instead. 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, @@ -1073,6 +1404,10 @@ 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*``. + If ``None``, copy all input arguments. + stage_out_argnames: Names of output arguments that need to be copied with ``GraphMode.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()``. module_preload_mode: Specify the devices where the module should be preloaded. @@ -1087,15 +1422,6 @@ def jax_callable( check_jax_version() - if graph_compatible is not None: - wp._src.utils.warn( - "The `graph_compatible` argument is deprecated, use `graph_mode` instead.", - DeprecationWarning, - stacklevel=3, - ) - if graph_compatible is False: - graph_mode = GraphMode.NONE - if graph_cache_max is None: graph_cache_max = FfiCallable.default_graph_cache_max @@ -1119,6 +1445,8 @@ def jax_callable( vmap_method, output_dims, in_out_argnames, + stage_in_argnames, + stage_out_argnames, graph_cache_max, module_preload_mode, ) diff --git a/mjx/pyproject.toml b/mjx/pyproject.toml index 3a025920..13ca80fa 100644 --- a/mjx/pyproject.toml +++ b/mjx/pyproject.toml @@ -36,7 +36,7 @@ dependencies = [ [project.optional-dependencies] warp = [ - "warp-lang==1.10.0", + "warp-lang==1.11.0", ] [project.scripts]