Import warp-lang 1.10.0.

PiperOrigin-RevId: 828583884
Change-Id: Ia1c6d8c41fa0785360b85287cab562228fcc80cf
This commit is contained in:
Baruch Tabanpour
2025-11-05 12:34:26 -08:00
committed by Copybara-Service
parent 3577c98f82
commit c34ac712a0
11 changed files with 568 additions and 205 deletions
+5
View File
@@ -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)
-----------------------------------
+5 -5
View File
@@ -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
@@ -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
@@ -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
@@ -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(
@@ -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,
-1
View File
@@ -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
+1 -1
View File
@@ -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, ...]):
-2
View File
@@ -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
+1 -1
View File
@@ -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',
+1 -1
View File
@@ -36,7 +36,7 @@ dependencies = [
[project.optional-dependencies]
warp = [
"warp-lang==1.9.1",
"warp-lang==1.10.0",
]
[project.scripts]