Import warp-lang 1.10.0.
PiperOrigin-RevId: 828583884 Change-Id: Ia1c6d8c41fa0785360b85287cab562228fcc80cf
This commit is contained in:
committed by
Copybara-Service
parent
3577c98f82
commit
c34ac712a0
@@ -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)
|
||||
-----------------------------------
|
||||
|
||||
|
||||
@@ -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
|
||||
|
||||
-2
@@ -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
|
||||
+10
-7
@@ -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
|
||||
+511
-185
@@ -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(
|
||||
+34
@@ -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,
|
||||
@@ -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
|
||||
|
||||
@@ -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, ...]):
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
@@ -36,7 +36,7 @@ dependencies = [
|
||||
|
||||
[project.optional-dependencies]
|
||||
warp = [
|
||||
"warp-lang==1.9.1",
|
||||
"warp-lang==1.10.0",
|
||||
]
|
||||
|
||||
[project.scripts]
|
||||
|
||||
Reference in New Issue
Block a user