Import Warp v1.16.0

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