From c34ac712a0dff19e180a725254b6350c01e62df9 Mon Sep 17 00:00:00 2001 From: Baruch Tabanpour Date: Wed, 5 Nov 2025 12:34:26 -0800 Subject: [PATCH] Import warp-lang 1.10.0. PiperOrigin-RevId: 828583884 Change-Id: Ia1c6d8c41fa0785360b85287cab562228fcc80cf --- doc/changelog.rst | 5 + mjx/cuda_requirements.txt | 10 +- .../{ => _src}/jax_experimental/__init__.py | 2 - .../jax_experimental/custom_call.py | 17 +- .../warp/{ => _src}/jax_experimental/ffi.py | 696 +++++++++++++----- .../{ => _src}/jax_experimental/xla_ffi.py | 34 + mjx/mujoco/mjx/warp/collision_driver.py | 1 - mjx/mujoco/mjx/warp/ffi.py | 2 +- mjx/mujoco/mjx/warp/forward.py | 2 - mjx/mujoco/mjx/warp/testspeed.py | 2 +- mjx/pyproject.toml | 2 +- 11 files changed, 568 insertions(+), 205 deletions(-) rename mjx/mujoco/mjx/third_party/warp/{ => _src}/jax_experimental/__init__.py (94%) rename mjx/mujoco/mjx/third_party/warp/{ => _src}/jax_experimental/custom_call.py (96%) rename mjx/mujoco/mjx/third_party/warp/{ => _src}/jax_experimental/ffi.py (56%) rename mjx/mujoco/mjx/third_party/warp/{ => _src}/jax_experimental/xla_ffi.py (93%) diff --git a/doc/changelog.rst b/doc/changelog.rst index cb09eddf..e26eab33 100644 --- a/doc/changelog.rst +++ b/doc/changelog.rst @@ -24,6 +24,11 @@ General - ``coef -> real`` - ``divisor -> real`` +MJX +^^^ + +- ``warp-lang`` optional dependency is updated to 1.10.0. ``pmap`` now works with MuJoCo Warp from MJX. + Version 3.3.7 (October 13, 2025) ----------------------------------- diff --git a/mjx/cuda_requirements.txt b/mjx/cuda_requirements.txt index afc63fe5..a540879c 100644 --- a/mjx/cuda_requirements.txt +++ b/mjx/cuda_requirements.txt @@ -16,8 +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.9.1 \ - --hash=sha256:3ce61319a268f38d135f4a26c4e75c29488a88cf3af273896e1fb2bf174131ec \ - --hash=sha256:2211e0cd9332ca44193801285cd1293c0b7580932da8da32e18d70b832adac12 \ - --hash=sha256:e58d0503b34d484d4699077ac4161a342d18714e5eb9bf63ede9d5df2188ecf0 \ - --hash=sha256:c971ed996004f729ceaf38b660831e95c837758df8ce6413a5d87d76652b0f88 +warp-lang==1.10.0 \ + --hash=sha256:428a6388ba8c9b3ded973226ddb5be59b16e3b3f28ac939a5036ad2d7cdc79ed \ + --hash=sha256:1bdff31e170b00c89fb9d8b647e906fefdcdcf8741925126d1c0be42783174fa \ + --hash=sha256:4aa8eb63cae5ee0d6dbdbdfc305d124140ec975475c97d4458f412396fb39eab \ + --hash=sha256:81f73055e76a6a3f2284cf2b5fe542a341156861101df9f291bd5b59925ff6e5 diff --git a/mjx/mujoco/mjx/third_party/warp/jax_experimental/__init__.py b/mjx/mujoco/mjx/third_party/warp/_src/jax_experimental/__init__.py similarity index 94% rename from mjx/mujoco/mjx/third_party/warp/jax_experimental/__init__.py rename to mjx/mujoco/mjx/third_party/warp/_src/jax_experimental/__init__.py index 89044207..3159bfe6 100644 --- a/mjx/mujoco/mjx/third_party/warp/jax_experimental/__init__.py +++ b/mjx/mujoco/mjx/third_party/warp/_src/jax_experimental/__init__.py @@ -12,5 +12,3 @@ # WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. # See the License for the specific language governing permissions and # limitations under the License. - -from .custom_call import jax_kernel diff --git a/mjx/mujoco/mjx/third_party/warp/jax_experimental/custom_call.py b/mjx/mujoco/mjx/third_party/warp/_src/jax_experimental/custom_call.py similarity index 96% rename from mjx/mujoco/mjx/third_party/warp/jax_experimental/custom_call.py rename to mjx/mujoco/mjx/third_party/warp/_src/jax_experimental/custom_call.py index ff8cc0c7..cd376477 100644 --- a/mjx/mujoco/mjx/third_party/warp/jax_experimental/custom_call.py +++ b/mjx/mujoco/mjx/third_party/warp/_src/jax_experimental/custom_call.py @@ -16,10 +16,12 @@ import ctypes 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 +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, strides_from_shape +from warp._src.utils import warn + +_wp_module_name_ = "warp.jax_experimental.custom_call" _jax_warp_p = None @@ -65,7 +67,8 @@ def jax_kernel(kernel, launch_dims=None, quiet=False): 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().", + "As of Warp release 1.10, the FFI version is the default implementation of jax_kernel(). " + "Pass quiet=True to disable this warning.", DeprecationWarning, ) @@ -130,7 +133,7 @@ def _warp_custom_callback(stream, buffers, opaque, opaque_len): assert hooks.forward, "Failed to find kernel entry point" # Launch the kernel. - wp.context.runtime.core.wp_cuda_launch_kernel( + wp._src.context.runtime.core.wp_cuda_launch_kernel( device.context, hooks.forward, bounds.size, 0, 256, hooks.forward_smem_bytes, kernel_params, stream ) @@ -142,7 +145,7 @@ def _create_jax_warp_primitive(): from jax._src.interpreters import batching from jax.interpreters import mlir from jax.interpreters.mlir import ir - from jax.google.hlo_helpers import custom_call + from jaxlib.hlo_helpers import custom_call global _jax_warp_p global _cc_callback diff --git a/mjx/mujoco/mjx/third_party/warp/jax_experimental/ffi.py b/mjx/mujoco/mjx/third_party/warp/_src/jax_experimental/ffi.py similarity index 56% rename from mjx/mujoco/mjx/third_party/warp/jax_experimental/ffi.py rename to mjx/mujoco/mjx/third_party/warp/_src/jax_experimental/ffi.py index b31d1a7b..64592ac0 100644 --- a/mjx/mujoco/mjx/third_party/warp/jax_experimental/ffi.py +++ b/mjx/mujoco/mjx/third_party/warp/_src/jax_experimental/ffi.py @@ -15,6 +15,7 @@ import collections import ctypes +import inspect import threading import traceback from enum import IntEnum @@ -23,17 +24,26 @@ from typing import Callable, Optional import jax import warp as wp -from warp.codegen import get_full_arg_spec, make_full_qualified_name -from warp.jax import get_jax_device -from warp.types import array_t, launch_bounds_t, strides_from_shape, type_to_warp +from warp._src.codegen import get_full_arg_spec, make_full_qualified_name +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 .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``. -""" +_wp_module_name_ = "warp.jax_experimental.ffi" + +# Type alias for differentiable kernel cache key +DiffKernelCacheKey = tuple[Callable, tuple, int, str, tuple[str, ...]] + +# Holders for the custom callbacks to keep them alive. +_FFI_KERNEL_REGISTRY: dict[str, "FfiKernel"] = {} +_FFI_DIFF_KERNEL_REGISTRY: dict[DiffKernelCacheKey, Callable] = {} +_FFI_CALLABLE_REGISTRY: dict[str, "FfiCallable"] = {} +_FFI_CALLBACK_REGISTRY: dict[str, ctypes.CFUNCTYPE] = {} +_FFI_REGISTRY_LOCK = threading.Lock() + +# Lock when XLA invokes callbacks from multiple threads. +_FFI_CALLBACK_LOCK = threading.Lock() def check_jax_version(): @@ -54,6 +64,12 @@ class GraphMode(IntEnum): WARP = 2 # let Warp capture a graph +class ModulePreloadMode(IntEnum): + NONE = 0 # don't preload modules + CURRENT_DEVICE = 1 # preload on currently active device + ALL_DEVICES = 2 # preload on all supported devices + + class FfiArg: def __init__(self, name, type, in_out=False): self.name = name @@ -67,7 +83,7 @@ class FfiArg: self.dtype_ndim = len(self.dtype_shape) self.jax_scalar_type = wp.dtype_to_jax(type.dtype._wp_scalar_type_) self.jax_ndim = type.ndim + self.dtype_ndim - elif type.dtype in wp.types.value_types: + elif type.dtype in wp._src.types.value_types: self.dtype_ndim = 0 self.dtype_shape = () self.jax_scalar_type = wp.dtype_to_jax(type.dtype) @@ -75,7 +91,7 @@ class FfiArg: else: raise TypeError(f"Invalid data type for array argument '{name}', expected scalar, vector, or matrix") self.warp_ndim = type.ndim - elif type in wp.types.value_types: + elif type in wp._src.types.value_types: self.dtype_ndim = 0 self.dtype_shape = () self.jax_scalar_type = wp.dtype_to_jax(type_to_warp(type)) @@ -92,13 +108,16 @@ class FfiLaunchDesc: class FfiKernel: - def __init__(self, kernel, num_outputs, vmap_method, launch_dims, output_dims, in_out_argnames): + def __init__( + self, kernel, num_outputs, vmap_method, launch_dims, output_dims, in_out_argnames, module_preload_mode + ): self.kernel = kernel self.name = generate_unique_name(kernel.func) self.num_outputs = num_outputs self.vmap_method = vmap_method self.launch_dims = launch_dims self.output_dims = output_dims + self.module_preload_mode = module_preload_mode self.first_array_arg = None self.launch_id = 0 self.launch_descriptors = {} @@ -251,9 +270,20 @@ class FfiKernel: input_output_aliases=self.input_output_aliases, ) - # ensure the kernel module is loaded before the callback, otherwise graph capture may fail - device = wp.device_from_jax(get_jax_device()) - self.kernel.module.load(device) + # preload on the specified devices + if self.module_preload_mode == ModulePreloadMode.CURRENT_DEVICE: + device = wp.device_from_jax(get_jax_device()) + self.kernel.module.load(device) + elif self.module_preload_mode == ModulePreloadMode.ALL_DEVICES: + for d in jax.local_devices(): + try: + dev = wp.device_from_jax(d) + 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 @@ -280,72 +310,75 @@ class FfiKernel: ) return None - # retrieve call info - attrs = decode_attrs(call_frame.contents.attrs) - launch_id = int(attrs["launch_id"]) - launch_desc = self.launch_descriptors[launch_id] + # Lock is required to prevent race conditions when callback is invoked + # from multiple threads, like with pmap. + with _FFI_CALLBACK_LOCK: + # retrieve call info + attrs = decode_attrs(call_frame.contents.attrs) + launch_id = int(attrs["launch_id"]) + launch_desc = self.launch_descriptors[launch_id] - num_inputs = call_frame.contents.args.size - inputs = ctypes.cast(call_frame.contents.args.args, ctypes.POINTER(ctypes.POINTER(XLA_FFI_Buffer))) + num_inputs = call_frame.contents.args.size + inputs = ctypes.cast(call_frame.contents.args.args, ctypes.POINTER(ctypes.POINTER(XLA_FFI_Buffer))) - num_outputs = call_frame.contents.rets.size - outputs = ctypes.cast(call_frame.contents.rets.rets, ctypes.POINTER(ctypes.POINTER(XLA_FFI_Buffer))) + num_outputs = call_frame.contents.rets.size + outputs = ctypes.cast(call_frame.contents.rets.rets, ctypes.POINTER(ctypes.POINTER(XLA_FFI_Buffer))) - assert num_inputs == self.num_inputs - assert num_outputs == self.num_outputs + assert num_inputs == self.num_inputs + assert num_outputs == self.num_outputs - launch_bounds = launch_bounds_t(launch_desc.launch_dims) + launch_bounds = launch_bounds_t(launch_desc.launch_dims) - # first kernel param is the launch bounds - kernel_params = (ctypes.c_void_p * (1 + self.num_kernel_args))() - kernel_params[0] = ctypes.addressof(launch_bounds) + # first kernel param is the launch bounds + kernel_params = (ctypes.c_void_p * (1 + self.num_kernel_args))() + kernel_params[0] = ctypes.addressof(launch_bounds) - arg_refs = [] + arg_refs = [] - # input and in-out args - for i, input_arg in enumerate(self.input_args): - if input_arg.is_array: - buffer = inputs[i].contents - shape = buffer.dims[: input_arg.type.ndim] - 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) + # input and in-out args + for i, input_arg in enumerate(self.input_args): + if input_arg.is_array: + buffer = inputs[i].contents + shape = buffer.dims[: input_arg.type.ndim] + 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) + for i, output_arg in enumerate(self.output_args): + buffer = outputs[i + self.num_in_out].contents + shape = buffer.dims[: output_arg.type.ndim] + 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 - # pure output args (skip in-out FFI buffers) - for i, output_arg in enumerate(self.output_args): - buffer = outputs[i + self.num_in_out].contents - shape = buffer.dims[: output_arg.type.ndim] - 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 + # get device and stream + device = wp.get_cuda_device(get_device_ordinal_from_callframe(call_frame.contents)) + stream = get_stream_from_callframe(call_frame.contents) - # get device and stream - device = wp.device_from_jax(get_jax_device()) - stream = get_stream_from_callframe(call_frame.contents) + # get kernel hooks + hooks = self.kernel.module.get_kernel_hooks(self.kernel, device) + assert hooks.forward, "Failed to find kernel entry point" - # get kernel hooks - hooks = self.kernel.module.get_kernel_hooks(self.kernel, device) - assert hooks.forward, "Failed to find kernel entry point" - - # launch the kernel - wp.context.runtime.core.wp_cuda_launch_kernel( - device.context, - hooks.forward, - launch_bounds.size, - 0, - 256, - hooks.forward_smem_bytes, - kernel_params, - stream, - ) + # launch the kernel + wp._src.context.runtime.core.wp_cuda_launch_kernel( + device.context, + hooks.forward, + launch_bounds.size, + 0, + 256, + hooks.forward_smem_bytes, + kernel_params, + stream, + ) except Exception as e: print(traceback.format_exc()) @@ -360,13 +393,26 @@ class FfiCallDesc: class FfiCallable: - def __init__(self, func, num_outputs, graph_mode, vmap_method, output_dims, in_out_argnames, graph_cache_max): + default_graph_cache_max: int | None = 32 + + def __init__( + self, + func, + num_outputs, + graph_mode, + vmap_method, + output_dims, + in_out_argnames, + graph_cache_max, + module_preload_mode, + ): self.func = func self.name = generate_unique_name(func) self.num_outputs = num_outputs self.vmap_method = vmap_method self.graph_mode = graph_mode self.output_dims = output_dims + self.module_preload_mode = module_preload_mode self.first_array_arg = None self.call_id = 0 self.call_descriptors = {} @@ -526,11 +572,22 @@ class FfiCallable: # has_side_effect=True, # force this function to execute even if outputs aren't used ) - # load the module + # preload on the specified devices # NOTE: if the target function uses kernels from different modules, they will not be loaded here - device = wp.device_from_jax(get_jax_device()) module = wp.get_module(self.func.__module__) - module.load(device) + if self.module_preload_mode == ModulePreloadMode.CURRENT_DEVICE: + device = wp.device_from_jax(get_jax_device()) + module.load(device) + elif self.module_preload_mode == ModulePreloadMode.ALL_DEVICES: + for d in jax.local_devices(): + try: + dev = wp.device_from_jax(d) + except Exception: + # ignore unsupported devices like TPUs + pass + # we only support CUDA devices for now + if dev.is_cuda: + module.load(dev) # save call data to be retrieved by callback call_id = self.call_id @@ -557,101 +614,105 @@ class FfiCallable: ) return None - # retrieve call info - # NOTE: this assumes that there's only one attribute - call_id (int64). - # A more general but slower approach is this: - # attrs = decode_attrs(call_frame.contents.attrs) - # call_id = int(attrs["call_id"]) - attr = ctypes.cast(call_frame.contents.attrs.attrs[0], ctypes.POINTER(XLA_FFI_Scalar)).contents - call_id = ctypes.cast(attr.value, ctypes.POINTER(ctypes.c_int64)).contents.value - call_desc = self.call_descriptors[call_id] + # Lock is required to prevent race conditions when callback is invoked + # from multiple threads, like with pmap. + with _FFI_CALLBACK_LOCK: + # retrieve call info + # NOTE: this assumes that there's only one attribute - call_id (int64). + # A more general but slower approach is this: + # attrs = decode_attrs(call_frame.contents.attrs) + # call_id = int(attrs["call_id"]) + attr = ctypes.cast(call_frame.contents.attrs.attrs[0], ctypes.POINTER(XLA_FFI_Scalar)).contents + call_id = ctypes.cast(attr.value, ctypes.POINTER(ctypes.c_int64)).contents.value + call_desc = self.call_descriptors[call_id] - num_inputs = call_frame.contents.args.size - inputs = ctypes.cast(call_frame.contents.args.args, ctypes.POINTER(ctypes.POINTER(XLA_FFI_Buffer))) + num_inputs = call_frame.contents.args.size + inputs = ctypes.cast(call_frame.contents.args.args, ctypes.POINTER(ctypes.POINTER(XLA_FFI_Buffer))) - num_outputs = call_frame.contents.rets.size - outputs = ctypes.cast(call_frame.contents.rets.rets, ctypes.POINTER(ctypes.POINTER(XLA_FFI_Buffer))) + num_outputs = call_frame.contents.rets.size + outputs = ctypes.cast(call_frame.contents.rets.rets, ctypes.POINTER(ctypes.POINTER(XLA_FFI_Buffer))) - assert num_inputs == self.num_inputs - assert num_outputs == self.num_outputs + assert num_inputs == self.num_inputs + assert num_outputs == self.num_outputs - cuda_stream = get_stream_from_callframe(call_frame.contents) + cuda_stream = get_stream_from_callframe(call_frame.contents) - if self.graph_mode == GraphMode.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] - capture_key = hash((call_id, *ip, *op)) - capture = self.captures.get(capture_key) + if self.graph_mode == GraphMode.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] + capture_key = hash((call_id, *ip, *op)) + capture = self.captures.get(capture_key) - # 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.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 + # 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.context.runtime.core.wp_cuda_graph_launch(graph.graph_exec, cuda_stream): - raise RuntimeError(f"Graph launch error: {wp.context.runtime.get_error_string()}") + 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()}") - # update the graph cache to keep recently used graphs alive - self.captures.move_to_end(capture_key) + # update the graph cache to keep recently used graphs alive + self.captures.move_to_end(capture_key) - # early out - return + # early out + return - device = wp.device_from_jax(get_jax_device()) - stream = wp.Stream(device, cuda_stream=cuda_stream) + 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) - # reconstruct the argument list - arg_list = [] + # 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 + # input and in-out args + for i, arg in enumerate(self.input_args): + if arg.is_array: + buffer = inputs[i].contents + shape = buffer.dims[: buffer.rank - arg.dtype_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 = buffer.dims[: buffer.rank - arg.dtype_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 = buffer.dims[: buffer.rank - arg.dtype_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 - with wp.ScopedStream(stream, sync_enter=False): - if stream.is_capturing: - # capturing with JAX - with wp.ScopedCapture(external=True) as capture: + # call the Python function with reconstructed arguments + 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) + # 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) + else: + # not capturing 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) - else: - # not capturing - self.func(*arg_list) except Exception as e: print(traceback.format_exc()) @@ -679,15 +740,15 @@ class FfiCallable: return len(self.captures) -# Holders for the custom callbacks to keep them alive. -_FFI_KERNEL_REGISTRY: dict[str, FfiKernel] = {} -_FFI_CALLABLE_REGISTRY: dict[str, FfiCallable] = {} -_FFI_CALLBACK_REGISTRY: dict[str, ctypes.CFUNCTYPE] = {} -_FFI_REGISTRY_LOCK = threading.Lock() - - def jax_kernel( - kernel, num_outputs=1, vmap_method="broadcast_all", launch_dims=None, output_dims=None, in_out_argnames=None + kernel, + num_outputs=1, + vmap_method="broadcast_all", + launch_dims=None, + output_dims=None, + in_out_argnames=None, + module_preload_mode=ModulePreloadMode.CURRENT_DEVICE, + enable_backward: bool = False, ): """Create a JAX callback from a Warp kernel. @@ -705,7 +766,12 @@ def jax_kernel( 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: Names of input-output arguments. + 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``. + Not supported when ``enable_backward=True``. + module_preload_mode: Specify the devices where the module should be preloaded. + enable_backward: Enable automatic differentiation for this kernel. Limitations: - All kernel arguments must be contiguous arrays or scalars. @@ -717,21 +783,257 @@ def jax_kernel( check_jax_version() - key = ( - kernel.func, - kernel.sig, - num_outputs, - vmap_method, - tuple(launch_dims) if launch_dims else launch_dims, - tuple(sorted(output_dims.items())) if output_dims else output_dims, + if not enable_backward: + key = ( + kernel.func, + kernel.sig, + num_outputs, + vmap_method, + tuple(launch_dims) if launch_dims else launch_dims, + tuple(sorted(output_dims.items())) if output_dims else output_dims, + module_preload_mode, + ) + + with _FFI_REGISTRY_LOCK: + if key not in _FFI_KERNEL_REGISTRY: + new_kernel = FfiKernel( + kernel, num_outputs, vmap_method, launch_dims, output_dims, in_out_argnames, module_preload_mode + ) + _FFI_KERNEL_REGISTRY[key] = new_kernel + + return _FFI_KERNEL_REGISTRY[key] + + # make sure the arguments are compatible with autodiff + if in_out_argnames: + raise NotImplementedError( + "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." + ) + + # Differentiable path: build a custom VJP wrapper inline. + # Infer the original kernel signature (names and annotations) + signature = inspect.signature(kernel.func) + + parameters = [p for p in signature.parameters.values() if p.kind == inspect.Parameter.POSITIONAL_OR_KEYWORD] + parameter_count = len(parameters) + num_inputs = parameter_count - num_outputs + + # determine static argument indices + static_args = [] + for i, p in enumerate(parameters[:num_inputs]): + param_type = p.annotation + if not isinstance(param_type, wp.array): + if param_type in wp._src.types.value_types: + static_args.append(i) + else: + raise TypeError(f"Invalid type for argument '{p.name}', expected array or scalar, got {type}") + + def _resolve_launch_dims(call_args): + # determine launch dimensions from the shape of the first input array + for i, p in enumerate(parameters[:num_inputs]): + param_type = p.annotation + if isinstance(param_type, wp.array): + arg = call_args[i] + arg_shape = tuple(arg.shape) + if hasattr(param_type.dtype, "_wp_scalar_type_"): + # vector/matrix array, trim trailing dimensions of JAX input array + return arg_shape[: param_type.ndim] + else: + # scalar array + return arg_shape + raise RuntimeError("Unable to determine launch dimensions, at least one input array is required") + + # 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:]) + + # update forward signature and annotations so jax_callable() sees a fully annotated function + fwd_kernel_wrapper.__signature__ = signature + fwd_kernel_wrapper.__annotations__ = {p.name: p.annotation for p in parameters} + fwd_kernel_wrapper.__annotations__["return"] = None + + jax_fwd_kernel = jax_callable(fwd_kernel_wrapper, num_outputs=num_outputs, vmap_method=vmap_method) + + # backward arguments only include static args once + bwd_arg_count = 2 * parameter_count - len(static_args) + + # Backward wrapper: launches adjoint with provided output gradients + def bwd_kernel_wrapper(*args): + if len(args) != bwd_arg_count: + raise RuntimeError(f"Invalid backward argument count, expected {bwd_arg_count} but got {len(args)}") + + inputs = list(args[:num_inputs]) + outputs = list(args[num_inputs:parameter_count]) + grad_out = list(args[parameter_count : parameter_count + num_outputs]) + grad_in = list(args[parameter_count + num_outputs :]) + + for i in static_args: + grad_in.insert(i, inputs[i]) + + for gi in grad_in: + if isinstance(gi, wp.array): + try: + gi.zero_() + except Exception as e: + wp.utils.warn(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. + wp.launch( + kernel, + dim=_resolve_launch_dims(inputs), + inputs=inputs, + outputs=outputs, + adj_inputs=grad_in, + adj_outputs=grad_out, + adjoint=True, + ) + + # Build the backward wrapper signature expected by jax_callable + bwd_input_params = parameters[:num_inputs] + bwd_output_params = parameters[num_inputs:parameter_count] + bwd_grad_output_params = [ + inspect.Parameter( + f"adj_{p.name}", + inspect.Parameter.POSITIONAL_OR_KEYWORD, + default=p.default, + annotation=p.annotation, + ) + for p in bwd_output_params + ] + + bwd_grad_input_params = [ + inspect.Parameter( + f"adj_{p.name}", + inspect.Parameter.POSITIONAL_OR_KEYWORD, + default=p.default, + annotation=p.annotation, + ) + for p in bwd_input_params + ] + for i in reversed(static_args): + del bwd_grad_input_params[i] + + # update backward signature and annotations so jax_callable() sees a fully annotated function + bwd_signature = bwd_input_params + bwd_output_params + bwd_grad_output_params + bwd_grad_input_params + bwd_kernel_wrapper.__signature__ = inspect.Signature(bwd_signature) + bwd_annotations = {} + for p in bwd_input_params: + bwd_annotations[p.name] = p.annotation + for p in bwd_output_params: + bwd_annotations[p.name] = p.annotation + for p in bwd_grad_output_params: + bwd_annotations[p.name] = p.annotation + for p in bwd_grad_input_params: + bwd_annotations[p.name] = p.annotation + bwd_annotations["return"] = None + bwd_kernel_wrapper.__annotations__ = bwd_annotations + + jax_bwd_kernel = jax_callable( + bwd_kernel_wrapper, + num_outputs=len(bwd_input_params) - len(static_args), + vmap_method=vmap_method, ) - with _FFI_REGISTRY_LOCK: - if key not in _FFI_KERNEL_REGISTRY: - new_kernel = FfiKernel(kernel, num_outputs, vmap_method, launch_dims, output_dims, in_out_argnames) - _FFI_KERNEL_REGISTRY[key] = new_kernel + differentiable_input_indices = [i for i in range(num_inputs) if i not in static_args] + differentiable_input_names = [parameters[i].name for i in differentiable_input_indices] - return _FFI_KERNEL_REGISTRY[key] + def fwd_function(*args): + outputs = jax_fwd_kernel(*args) + non_static_inputs = list(args) + for i in reversed(static_args): + del non_static_inputs[i] + # Normalize to tuple for consistent handling + if num_outputs == 1: + outputs_tuple = (outputs,) if not isinstance(outputs, (list, tuple)) else (outputs[0],) + else: + outputs_tuple = outputs if isinstance(outputs, tuple) else tuple(outputs) + return outputs, (tuple(non_static_inputs), outputs_tuple) + + def bwd_function(*bwd_args): + nondiff_vals = list(bwd_args[: len(static_args)]) + residuals = bwd_args[len(static_args)] + grad_out_args = bwd_args[len(static_args) + 1 :] + + non_static_inputs, output_vals_tuple = residuals + + input_vals = list(non_static_inputs) + for i, v in zip(static_args, nondiff_vals): + input_vals.insert(i, v) + + # Normalize grad outputs and handle nested containers (e.g., single tuple for multi-output) + if num_outputs == 1: + go = grad_out_args[0] + grad_out_tuple = tuple(go) if isinstance(go, (list, tuple)) else (go,) + else: + if len(grad_out_args) == 1 and isinstance(grad_out_args[0], (list, tuple)): + grad_out_tuple = tuple(grad_out_args[0]) + else: + grad_out_tuple = tuple(grad_out_args) + bwd_call_args = list(input_vals) + list(output_vals_tuple) + list(grad_out_tuple) + + out_dims_map = {} + param_ann = {p.name: p.annotation for p in parameters[:num_inputs]} + for name, val in zip(differentiable_input_names, non_static_inputs): + ann = param_ann.get(name) + if ann is None: + continue + # Check if annotation is a warp array type (annotation is an instance of wp.array) + is_array_ann = isinstance(ann, wp.array) + if not is_array_ann: + continue + dtype_ndim = 0 + # Extract dtype ndim if it's a vector/matrix type + if hasattr(ann, "dtype") and hasattr(ann.dtype, "_wp_scalar_type_"): + dtype_ndim = len(ann.dtype._shape_) + warp_ndim = getattr(ann, "ndim", 0) + vshape = tuple(val.shape) + if warp_ndim == 0: + continue + if dtype_ndim > 0: + core_rank = max(0, len(vshape) - dtype_ndim) + warp_dims = vshape[max(0, core_rank - warp_ndim) : core_rank] + else: + warp_dims = vshape[-warp_ndim:] + out_dims_map[f"adj_{name}"] = tuple(warp_dims) + + non_static_input_grads = jax_bwd_kernel(*bwd_call_args, output_dims=out_dims_map) + return tuple(non_static_input_grads) + + jax_func = jax.custom_vjp(jax_fwd_kernel, nondiff_argnums=tuple(static_args)) + jax_func.defvjp(fwd_function, bwd_function) + + if static_args: + static_names = [parameters[i].name for i in static_args] + + def _user_callable(*args): + return jax_func(*args) + + _user_callable.__signature__ = signature + + # Cache differentiable wrapper + key = (kernel.func, kernel.sig, num_outputs, vmap_method, 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] + + # Cache differentiable wrapper (no static args) + key = (kernel.func, kernel.sig, num_outputs, vmap_method, ()) + 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 def jax_callable( @@ -743,6 +1045,7 @@ def jax_callable( output_dims=None, in_out_argnames=None, graph_cache_max: int | None = None, + module_preload_mode: ModulePreloadMode = ModulePreloadMode.CURRENT_DEVICE, ): """Create a JAX callback from an annotated Python function. @@ -767,9 +1070,12 @@ def jax_callable( 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: Names of input-output arguments. + 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``. 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``. + 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. Limitations: - All kernel arguments must be contiguous arrays or scalars. @@ -782,7 +1088,7 @@ def jax_callable( check_jax_version() if graph_compatible is not None: - wp.utils.warn( + wp._src.utils.warn( "The `graph_compatible` argument is deprecated, use `graph_mode` instead.", DeprecationWarning, stacklevel=3, @@ -791,7 +1097,7 @@ def jax_callable( graph_mode = GraphMode.NONE if graph_cache_max is None: - graph_cache_max = jax_callable_default_graph_cache_max + graph_cache_max = FfiCallable.default_graph_cache_max # Note: we don't include graph_cache_max in the key, it is applied below. key = ( @@ -800,6 +1106,7 @@ def jax_callable( graph_mode, vmap_method, tuple(sorted(output_dims.items())) if output_dims else output_dims, + module_preload_mode, ) with _FFI_REGISTRY_LOCK: @@ -813,6 +1120,7 @@ def jax_callable( output_dims, in_out_argnames, graph_cache_max, + module_preload_mode, ) _FFI_CALLABLE_REGISTRY[key] = callable else: @@ -822,6 +1130,20 @@ def jax_callable( return 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``. + """ + 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``. + """ + FfiCallable.default_graph_cache_max = cache_max + + def clear_jax_callable_graph_cache(callable: FfiCallable | None = None): """Clear the graph cache of the given callable or all callables if ``None``.""" @@ -878,19 +1200,23 @@ def register_ffi_callback(name: str, func: Callable, graph_compatible: bool = Tr ) return None - attrs = decode_attrs(call_frame.contents.attrs) + # Lock is required to prevent race conditions when callback is invoked + # from multiple threads, like with pmap. + with _FFI_CALLBACK_LOCK: + attrs = decode_attrs(call_frame.contents.attrs) - input_count = call_frame.contents.args.size - inputs = ctypes.cast(call_frame.contents.args.args, ctypes.POINTER(ctypes.POINTER(XLA_FFI_Buffer))) - inputs = [FfiBuffer(inputs[i].contents) for i in range(input_count)] + input_count = call_frame.contents.args.size + inputs = ctypes.cast(call_frame.contents.args.args, ctypes.POINTER(ctypes.POINTER(XLA_FFI_Buffer))) + inputs = [FfiBuffer(inputs[i].contents) for i in range(input_count)] - output_count = call_frame.contents.rets.size - outputs = ctypes.cast(call_frame.contents.rets.rets, ctypes.POINTER(ctypes.POINTER(XLA_FFI_Buffer))) - outputs = [FfiBuffer(outputs[i].contents) for i in range(output_count)] + output_count = call_frame.contents.rets.size + outputs = ctypes.cast(call_frame.contents.rets.rets, ctypes.POINTER(ctypes.POINTER(XLA_FFI_Buffer))) + outputs = [FfiBuffer(outputs[i].contents) for i in range(output_count)] - ctx = ExecutionContext(call_frame.contents) + ctx = ExecutionContext(call_frame.contents) + + func(inputs, outputs, attrs, ctx) - func(inputs, outputs, attrs, ctx) except Exception as e: print(traceback.format_exc()) return create_ffi_error( diff --git a/mjx/mujoco/mjx/third_party/warp/jax_experimental/xla_ffi.py b/mjx/mujoco/mjx/third_party/warp/_src/jax_experimental/xla_ffi.py similarity index 93% rename from mjx/mujoco/mjx/third_party/warp/jax_experimental/xla_ffi.py rename to mjx/mujoco/mjx/third_party/warp/_src/jax_experimental/xla_ffi.py index 8e311143..2da12c4e 100644 --- a/mjx/mujoco/mjx/third_party/warp/jax_experimental/xla_ffi.py +++ b/mjx/mujoco/mjx/third_party/warp/_src/jax_experimental/xla_ffi.py @@ -21,6 +21,8 @@ import numpy as np import warp as wp +_wp_module_name_ = "warp.jax_experimental.xla_ffi" + ####################################################################### # ctypes structures and enums for XLA's FFI API: # https://github.com/openxla/xla/blob/a1a5e62fbffa3a3b6c409d72607456cf5b353a22/xla/ffi/api/c_api.h @@ -379,6 +381,24 @@ class XLA_FFI_Stream_Get_Args(ctypes.Structure): XLA_FFI_Stream_Get = ctypes.CFUNCTYPE(ctypes.c_void_p, ctypes.POINTER(XLA_FFI_Stream_Get_Args)) +# struct XLA_FFI_DeviceOrdinal_Get { +# size_t struct_size; +# XLA_FFI_Extension_Base* extension_start; +# XLA_FFI_ExecutionContext* ctx; +# int32_t device_ordinal; // out +# }; +class XLA_FFI_DeviceOrdinal_Get_Args(ctypes.Structure): + _fields_ = ( + ("struct_size", ctypes.c_size_t), + ("extension_start", ctypes.POINTER(XLA_FFI_Extension_Base)), + ("ctx", ctypes.c_void_p), # XLA_FFI_ExecutionContext* + ("device_ordinal", ctypes.c_int32), + ) # // out + + +XLA_FFI_DeviceOrdinal_Get = ctypes.CFUNCTYPE(ctypes.c_void_p, ctypes.POINTER(XLA_FFI_DeviceOrdinal_Get_Args)) + + # struct XLA_FFI_Api { # size_t struct_size; # XLA_FFI_Extension_Base* extension_start; @@ -402,6 +422,8 @@ XLA_FFI_Stream_Get = ctypes.CFUNCTYPE(ctypes.c_void_p, ctypes.POINTER(XLA_FFI_St # _XLA_FFI_API_STRUCT_FIELD(XLA_FFI_Future_Create); # _XLA_FFI_API_STRUCT_FIELD(XLA_FFI_Future_SetAvailable); # _XLA_FFI_API_STRUCT_FIELD(XLA_FFI_Future_SetError); +# _XLA_FFI_API_STRUCT_FIELD(XLA_FFI_RunId_Get); +# _XLA_FFI_API_STRUCT_FIELD(XLA_FFI_DeviceOrdinal_Get); # }; class XLA_FFI_Api(ctypes.Structure): _fields_ = ( @@ -425,6 +447,9 @@ class XLA_FFI_Api(ctypes.Structure): ("XLA_FFI_Future_Create", ctypes.c_void_p), # XLA_FFI_Future_Create ("XLA_FFI_Future_SetAvailable", ctypes.c_void_p), # XLA_FFI_Future_SetAvailable ("XLA_FFI_Future_SetError", ctypes.c_void_p), # XLA_FFI_Future_SetError + # TODO(chaserileyroberts): Make this return the correct value and not a c_void_p. + ("XLA_FFI_RunId_Get", ctypes.c_void_p), # XLA_FFI_RunId_Get + ("XLA_FFI_DeviceOrdinal_Get", XLA_FFI_DeviceOrdinal_Get), # XLA_FFI_DeviceOrdinal_Get ) @@ -570,6 +595,15 @@ def get_stream_from_callframe(call_frame): return get_stream_args.stream +def get_device_ordinal_from_callframe(call_frame): + api = call_frame.api + get_device_args = XLA_FFI_DeviceOrdinal_Get_Args( + ctypes.sizeof(XLA_FFI_DeviceOrdinal_Get_Args), ctypes.POINTER(XLA_FFI_Extension_Base)(), call_frame.ctx, 0 + ) + api.contents.XLA_FFI_DeviceOrdinal_Get(get_device_args) + return get_device_args.device_ordinal + + _dtype_from_ffi = { XLA_FFI_DataType.S8: wp.int8, XLA_FFI_DataType.S16: wp.int16, diff --git a/mjx/mujoco/mjx/warp/collision_driver.py b/mjx/mujoco/mjx/warp/collision_driver.py index 498e31bf..0675b0e5 100644 --- a/mjx/mujoco/mjx/warp/collision_driver.py +++ b/mjx/mujoco/mjx/warp/collision_driver.py @@ -42,7 +42,6 @@ _e = mjwarp.Constraint( **{f.name: None for f in dataclasses.fields(mjwarp.Constraint) 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 a04a232b..ac2040d2 100644 --- a/mjx/mujoco/mjx/warp/ffi.py +++ b/mjx/mujoco/mjx/warp/ffi.py @@ -25,7 +25,7 @@ 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.jax_experimental import ffi +from mujoco.mjx.third_party.warp._src.jax_experimental import ffi def flatten_signature(signature: inspect.Signature, args: Tuple[Any, ...]): diff --git a/mjx/mujoco/mjx/warp/forward.py b/mjx/mujoco/mjx/warp/forward.py index 45102b61..ae66d6e7 100644 --- a/mjx/mujoco/mjx/warp/forward.py +++ b/mjx/mujoco/mjx/warp/forward.py @@ -42,7 +42,6 @@ _e = mjwarp.Constraint( **{f.name: None for f in dataclasses.fields(mjwarp.Constraint) if f.init} ) - @ffi.format_args_for_warp def _forward_shim( # Model @@ -2065,7 +2064,6 @@ _e = mjwarp.Constraint( **{f.name: None for f in dataclasses.fields(mjwarp.Constraint) if f.init} ) - @ffi.format_args_for_warp def _step_shim( # Model diff --git a/mjx/mujoco/mjx/warp/testspeed.py b/mjx/mujoco/mjx/warp/testspeed.py index 3cf497ad..8faba97f 100644 --- a/mjx/mujoco/mjx/warp/testspeed.py +++ b/mjx/mujoco/mjx/warp/testspeed.py @@ -30,7 +30,7 @@ from mujoco.mjx.warp import forward as wp_forward from mujoco.mjx.warp import smooth as wp_smooth import mujoco.mjx.third_party.mujoco_warp as mjwarp import warp as wp -from mujoco.mjx.third_party.warp.jax_experimental import ffi as warp_ffi +from mujoco.mjx.third_party.warp._src.jax_experimental import ffi as warp_ffi _MODELFILE = flags.DEFINE_string( 'modelfile', diff --git a/mjx/pyproject.toml b/mjx/pyproject.toml index eea6902d..fff25206 100644 --- a/mjx/pyproject.toml +++ b/mjx/pyproject.toml @@ -36,7 +36,7 @@ dependencies = [ [project.optional-dependencies] warp = [ - "warp-lang==1.9.1", + "warp-lang==1.10.0", ] [project.scripts]