From fbb95d5e8c8785e1f8d212675d7abed66293c853 Mon Sep 17 00:00:00 2001 From: Erik Frey Date: Fri, 31 Oct 2025 11:46:10 -0700 Subject: [PATCH] Import NVIDIA/warp from GitHub. PiperOrigin-RevId: 826571491 Change-Id: I294abc5aa3714345ab295636d12428293ea69285 --- .../warp/jax_experimental/__init__.py | 15 +- .../warp/jax_experimental/custom_call.py | 373 +------ .../third_party/warp/jax_experimental/ffi.py | 955 +----------------- .../warp/jax_experimental/xla_ffi.py | 610 +---------- 4 files changed, 44 insertions(+), 1909 deletions(-) diff --git a/mjx/mujoco/mjx/third_party/warp/jax_experimental/__init__.py b/mjx/mujoco/mjx/third_party/warp/jax_experimental/__init__.py index 89044207..920cdb3f 100644 --- a/mjx/mujoco/mjx/third_party/warp/jax_experimental/__init__.py +++ b/mjx/mujoco/mjx/third_party/warp/jax_experimental/__init__.py @@ -13,4 +13,17 @@ # See the License for the specific language governing permissions and # limitations under the License. -from .custom_call import jax_kernel +# isort: skip_file + +from warp._src.jax_experimental.ffi import GraphMode as GraphMode +from warp._src.jax_experimental.ffi import jax_kernel as jax_kernel +from warp._src.jax_experimental.ffi import jax_callable as jax_callable +from warp._src.jax_experimental.ffi import register_ffi_callback as register_ffi_callback + +from warp._src.jax_experimental.ffi import ( + get_jax_callable_default_graph_cache_max as get_jax_callable_default_graph_cache_max, +) +from warp._src.jax_experimental.ffi import ( + set_jax_callable_default_graph_cache_max as set_jax_callable_default_graph_cache_max, +) +from warp._src.jax_experimental.ffi import clear_jax_callable_graph_cache as clear_jax_callable_graph_cache diff --git a/mjx/mujoco/mjx/third_party/warp/jax_experimental/custom_call.py b/mjx/mujoco/mjx/third_party/warp/jax_experimental/custom_call.py index ff8cc0c7..fa18e9c5 100644 --- a/mjx/mujoco/mjx/third_party/warp/jax_experimental/custom_call.py +++ b/mjx/mujoco/mjx/third_party/warp/jax_experimental/custom_call.py @@ -1,4 +1,4 @@ -# SPDX-FileCopyrightText: Copyright (c) 2024 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-FileCopyrightText: Copyright (c) 2025 NVIDIA CORPORATION & AFFILIATES. All rights reserved. # SPDX-License-Identifier: Apache-2.0 # # Licensed under the Apache License, Version 2.0 (the "License"); @@ -13,374 +13,17 @@ # See the License for the specific language governing permissions and # limitations under the License. -import ctypes +# isort: skip_file -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.jax_experimental.custom_call import jax_kernel as jax_kernel -_jax_warp_p = None -# Holder for the custom callback to keep it alive. -_cc_callback = None -_registered_kernels = [None] -_registered_kernel_to_id = {} +# TODO: Remove after cleaning up the public API. +from warp._src.jax_experimental import custom_call as _custom_call -def jax_kernel(kernel, launch_dims=None, quiet=False): - """Create a Jax primitive from a Warp kernel. - NOTE: This is an experimental feature under development. +def __getattr__(name): + from warp._src.utils import get_deprecated_api - Args: - kernel: The Warp kernel to be wrapped. - launch_dims: Optional. Specify the kernel launch dimensions. If None, - dimensions are inferred from the shape of the first argument. - This option when set will specify the output dimensions. - quiet: Optional. If True, suppress deprecation warnings with newer JAX versions. - - Limitations: - - All kernel arguments must be contiguous arrays. - - Input arguments are followed by output arguments in the Warp kernel definition. - - There must be at least one input argument and at least one output argument. - - Only the CUDA backend is supported. - """ - - import jax - - # check if JAX version supports this - if jax.__version_info__ < (0, 4, 25) or jax.__version_info__ >= (0, 8, 0): - msg = ( - "This version of jax_kernel() requires JAX version 0.4.25 - 0.7.x, " - f"but installed JAX version is {jax.__version_info__}." - ) - if jax.__version_info__ >= (0, 8, 0): - msg += " Please use warp.jax_experimental.ffi.jax_kernel instead." - raise RuntimeError(msg) - - # deprecation warning - if jax.__version_info__ >= (0, 5, 0) and not quiet: - 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().", - DeprecationWarning, - ) - - if _jax_warp_p is None: - # Create and register the primitive - _create_jax_warp_primitive() - if kernel not in _registered_kernel_to_id: - id = len(_registered_kernels) - _registered_kernels.append(kernel) - _registered_kernel_to_id[kernel] = id - else: - id = _registered_kernel_to_id[kernel] - - def bind(*args): - return _jax_warp_p.bind(*args, kernel=id, launch_dims=launch_dims) - - return bind - - -def _warp_custom_callback(stream, buffers, opaque, opaque_len): - # The descriptor is the form - # || - # Example: 42|16,32|16,32;100;16,32 - kernel_id_str, dim_str, args_str = opaque.decode().split("|") - - # Get the kernel from the registry. - kernel_id = int(kernel_id_str) - kernel = _registered_kernels[kernel_id] - - # Parse launch dimensions. - dims = [int(d) for d in dim_str.split(",")] - bounds = launch_bounds_t(dims) - - # Parse arguments. - arg_strings = args_str.split(";") - num_args = len(arg_strings) - assert num_args == len(kernel.adj.args), "Incorrect number of arguments" - - # First param is the launch bounds. - kernel_params = (ctypes.c_void_p * (1 + num_args))() - kernel_params[0] = ctypes.addressof(bounds) - - # Parse array descriptors. - args = [] - for i in range(num_args): - dtype = kernel.adj.args[i].type.dtype - shape = [int(d) for d in arg_strings[i].split(",")] - strides = strides_from_shape(shape, dtype) - - arr = array_t(buffers[i], 0, len(shape), shape, strides) - args.append(arr) # keep a reference - arg_ptr = ctypes.addressof(arr) - - kernel_params[i + 1] = arg_ptr - - # Get current device. - device = wp.device_from_jax(get_jax_device()) - - # Get kernel hooks. - # Note: module was loaded during jit lowering. - hooks = kernel.module.get_kernel_hooks(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, bounds.size, 0, 256, hooks.forward_smem_bytes, kernel_params, stream - ) - - -def _create_jax_warp_primitive(): - from functools import reduce - - import jax - from jax._src.interpreters import batching - from jax.interpreters import mlir - from jax.interpreters.mlir import ir - from jax.google.hlo_helpers import custom_call - - global _jax_warp_p - global _cc_callback - - # Create and register the primitive. - # TODO add default implementation that calls the kernel via warp. - try: - # newer JAX versions - import jax.extend - - _jax_warp_p = jax.extend.core.Primitive("jax_warp") - except (ImportError, AttributeError): - # older JAX versions - _jax_warp_p = jax.core.Primitive("jax_warp") - _jax_warp_p.multiple_results = True - - # TODO Just launch the kernel directly, but make sure the argument - # shapes are massaged the same way as below so that vmap works. - def impl(*args): - raise Exception("Not implemented") - - _jax_warp_p.def_impl(impl) - - # Auto-batching. Make sure all the arguments are fully broadcasted - # so that Warp is not confused about dimensions. - def vectorized_multi_batcher(args, dims, **params): - # Figure out the number of outputs. - wp_kernel = _registered_kernels[params["kernel"]] - output_count = len(wp_kernel.adj.args) - len(args) - shape, dim = next((a.shape, d) for a, d in zip(args, dims) if d is not None) - size = shape[dim] - args = [batching.bdim_at_front(a, d, size) if len(a.shape) else a for a, d in zip(args, dims)] - # Create the batched primitive. - return _jax_warp_p.bind(*args, **params), [dims[0]] * output_count - - batching.primitive_batchers[_jax_warp_p] = vectorized_multi_batcher - - def get_vecmat_shape(warp_type): - if hasattr(warp_type.dtype, "_shape_"): - return warp_type.dtype._shape_ - return [] - - def strip_vecmat_dimensions(warp_arg, actual_shape): - shape = get_vecmat_shape(warp_arg.type) - for i, s in enumerate(reversed(shape)): - item = actual_shape[-i - 1] - if s != item: - raise Exception(f"The vector/matrix shape for argument {warp_arg.label} does not match") - return actual_shape[: len(actual_shape) - len(shape)] - - def collapse_into_leading_dimension(warp_arg, actual_shape): - if len(actual_shape) < warp_arg.type.ndim: - raise Exception(f"Argument {warp_arg.label} has too few non-matrix/vector dimensions") - index_rest = len(actual_shape) - warp_arg.type.ndim + 1 - leading_size = reduce(lambda x, y: x * y, actual_shape[:index_rest]) - return [leading_size] + actual_shape[index_rest:] - - # Infer array dimensions from input type. - def infer_dimensions(warp_arg, actual_shape): - actual_shape = strip_vecmat_dimensions(warp_arg, actual_shape) - return collapse_into_leading_dimension(warp_arg, actual_shape) - - def base_type_to_jax(warp_dtype): - if hasattr(warp_dtype, "_wp_scalar_type_"): - return wp.dtype_to_jax(warp_dtype._wp_scalar_type_) - return wp.dtype_to_jax(warp_dtype) - - def base_type_to_jax_ir(warp_dtype): - warp_to_jax_dict = { - wp.float16: ir.F16Type.get(), - wp.float32: ir.F32Type.get(), - wp.float64: ir.F64Type.get(), - wp.int8: ir.IntegerType.get_signless(8), - wp.int16: ir.IntegerType.get_signless(16), - wp.int32: ir.IntegerType.get_signless(32), - wp.int64: ir.IntegerType.get_signless(64), - wp.uint8: ir.IntegerType.get_unsigned(8), - wp.uint16: ir.IntegerType.get_unsigned(16), - wp.uint32: ir.IntegerType.get_unsigned(32), - wp.uint64: ir.IntegerType.get_unsigned(64), - } - if hasattr(warp_dtype, "_wp_scalar_type_"): - warp_dtype = warp_dtype._wp_scalar_type_ - jax_dtype = warp_to_jax_dict.get(warp_dtype) - if jax_dtype is None: - raise TypeError(f"Invalid or unsupported data type: {warp_dtype}") - return jax_dtype - - def base_type_is_compatible(warp_type, jax_ir_type): - jax_ir_to_warp = { - "f16": wp.float16, - "f32": wp.float32, - "f64": wp.float64, - "i8": wp.int8, - "i16": wp.int16, - "i32": wp.int32, - "i64": wp.int64, - "ui8": wp.uint8, - "ui16": wp.uint16, - "ui32": wp.uint32, - "ui64": wp.uint64, - } - expected_warp_type = jax_ir_to_warp.get(str(jax_ir_type)) - if expected_warp_type is not None: - if hasattr(warp_type, "_wp_scalar_type_"): - return warp_type._wp_scalar_type_ == expected_warp_type - else: - return warp_type == expected_warp_type - else: - raise TypeError(f"Invalid or unsupported data type: {jax_ir_type}") - - # Abstract evaluation. - def jax_warp_abstract(*args, kernel=None, launch_dims=None): - wp_kernel = _registered_kernels[kernel] - # All the extra arguments to the warp kernel are outputs. - warp_outputs = [o.type for o in wp_kernel.adj.args[len(args) :]] - - if launch_dims is None: - # Use the first input dimension to infer the output's dimensions if launch_dims is not provided - dims = strip_vecmat_dimensions(wp_kernel.adj.args[0], list(args[0].shape)) - else: - dims = launch_dims - - jax_outputs = [] - for o in warp_outputs: - shape = list(dims) + list(get_vecmat_shape(o)) - dtype = base_type_to_jax(o.dtype) - jax_outputs.append(jax.core.ShapedArray(shape, dtype)) - return jax_outputs - - _jax_warp_p.def_abstract_eval(jax_warp_abstract) - - # Lowering to MLIR. - - # Create python-land custom call target. - CCALLFUNC = ctypes.CFUNCTYPE( - ctypes.c_voidp, ctypes.c_void_p, ctypes.POINTER(ctypes.c_void_p), ctypes.c_char_p, ctypes.c_size_t - ) - _cc_callback = CCALLFUNC(_warp_custom_callback) - ccall_address = ctypes.cast(_cc_callback, ctypes.c_void_p) - - # Put the custom call into a capsule, as required by XLA. - PyCapsule_Destructor = ctypes.CFUNCTYPE(None, ctypes.py_object) - PyCapsule_New = ctypes.pythonapi.PyCapsule_New - PyCapsule_New.restype = ctypes.py_object - PyCapsule_New.argtypes = (ctypes.c_void_p, ctypes.c_char_p, PyCapsule_Destructor) - capsule = PyCapsule_New(ccall_address.value, b"xla._CUSTOM_CALL_TARGET", PyCapsule_Destructor(0)) - - # Register the callback in XLA. - try: - # newer JAX versions - jax.ffi.register_ffi_target("warp_call", capsule, platform="gpu", api_version=0) - except AttributeError: - # older JAX versions - jax.lib.xla_client.register_custom_call_target("warp_call", capsule, platform="gpu") - - def default_layout(shape): - return range(len(shape) - 1, -1, -1) - - def warp_call_lowering(ctx, *args, kernel=None, launch_dims=None): - if not kernel: - raise Exception("Unknown kernel id " + str(kernel)) - wp_kernel = _registered_kernels[kernel] - - # TODO This may not be necessary, but it is perhaps better not to be - # mucking with kernel loading while already running the workload. - module = wp_kernel.module - device = wp.device_from_jax(get_jax_device()) - if not module.load(device): - raise Exception("Could not load kernel on device") - - if launch_dims is None: - # Infer dimensions from the first input. - warp_arg0 = wp_kernel.adj.args[0] - actual_shape0 = ir.RankedTensorType(args[0].type).shape - dims = strip_vecmat_dimensions(warp_arg0, actual_shape0) - warp_dims = collapse_into_leading_dimension(warp_arg0, dims) - else: - dims = launch_dims - warp_dims = launch_dims - # Figure out the types and shapes of the input arrays. - arg_strings = [] - operand_layouts = [] - for actual, warg in zip(args, wp_kernel.adj.args): - wtype = warg.type - rtt = ir.RankedTensorType(actual.type) - - if not isinstance(wtype, wp.array): - raise Exception("Only contiguous arrays are supported for Jax kernel arguments") - - if not base_type_is_compatible(wtype.dtype, rtt.element_type): - raise TypeError( - f"Incompatible data type for argument '{warg.label}', expected {type_str(wtype.dtype)}, got {rtt.element_type}" - ) - - # Infer array dimension (by removing the vector/matrix dimensions and - # collapsing the initial dimensions). - shape = infer_dimensions(warg, rtt.shape) - - if len(shape) != wtype.ndim: - raise TypeError(f"Incompatible array dimensionality for argument '{warg.label}'") - - arg_strings.append(",".join([str(d) for d in shape])) - operand_layouts.append(default_layout(rtt.shape)) - - # Figure out the types and shapes of the output arrays. - result_types = [] - result_layouts = [] - for warg in wp_kernel.adj.args[len(args) :]: - wtype = warg.type - - if not isinstance(wtype, wp.array): - raise Exception("Only contiguous arrays are supported for Jax kernel arguments") - - # Infer dimensions from the first input. - arg_strings.append(",".join([str(d) for d in warp_dims])) - - result_shape = list(dims) + list(get_vecmat_shape(wtype)) - result_types.append(ir.RankedTensorType.get(result_shape, base_type_to_jax_ir(wtype.dtype))) - result_layouts.append(default_layout(result_shape)) - - # Build opaque descriptor for callback. - shape_str = ",".join([str(d) for d in warp_dims]) - args_str = ";".join(arg_strings) - descriptor = f"{kernel}|{shape_str}|{args_str}" - - out = custom_call( - b"warp_call", - result_types=result_types, - operands=args, - backend_config=descriptor.encode("utf-8"), - operand_layouts=operand_layouts, - result_layouts=result_layouts, - ).results - return out - - mlir.register_lowering( - _jax_warp_p, - warp_call_lowering, - platform="gpu", - ) + return get_deprecated_api(_custom_call, "wp.jax_experimental", name) diff --git a/mjx/mujoco/mjx/third_party/warp/jax_experimental/ffi.py b/mjx/mujoco/mjx/third_party/warp/jax_experimental/ffi.py index b31d1a7b..164ab94d 100644 --- a/mjx/mujoco/mjx/third_party/warp/jax_experimental/ffi.py +++ b/mjx/mujoco/mjx/third_party/warp/jax_experimental/ffi.py @@ -13,948 +13,27 @@ # See the License for the specific language governing permissions and # limitations under the License. -import collections -import ctypes -import threading -import traceback -from enum import IntEnum -from typing import Callable, Optional +# isort: skip_file -import jax +from warp._src.jax_experimental.ffi import GraphMode as GraphMode +from warp._src.jax_experimental.ffi import jax_kernel as jax_kernel +from warp._src.jax_experimental.ffi import jax_callable as jax_callable +from warp._src.jax_experimental.ffi import register_ffi_callback as register_ffi_callback -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.jax_experimental.ffi import ( + get_jax_callable_default_graph_cache_max as get_jax_callable_default_graph_cache_max, +) +from warp._src.jax_experimental.ffi import ( + set_jax_callable_default_graph_cache_max as set_jax_callable_default_graph_cache_max, +) +from warp._src.jax_experimental.ffi import clear_jax_callable_graph_cache as clear_jax_callable_graph_cache -from .xla_ffi import * +# TODO: Remove after cleaning up the public API. -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``. -""" +from warp._src.jax_experimental import ffi as _ffi -def check_jax_version(): - # check if JAX version supports this - if jax.__version_info__ < (0, 5, 0): - msg = ( - "This version of jax_kernel() requires JAX version 0.5.0 or higher, " - f"but installed JAX version is {jax.__version_info__}." - ) - if jax.__version_info__ >= (0, 4, 25): - msg += " Please use warp.jax_experimental.custom_call.jax_kernel instead." - raise RuntimeError(msg) +def __getattr__(name): + from warp._src.utils import get_deprecated_api - -class GraphMode(IntEnum): - NONE = 0 # don't capture a graph - JAX = 1 # let JAX capture a graph - WARP = 2 # let Warp capture a graph - - -class FfiArg: - def __init__(self, name, type, in_out=False): - self.name = name - self.type = type - self.in_out = in_out - self.is_array = isinstance(type, wp.array) - - if self.is_array: - if hasattr(type.dtype, "_wp_scalar_type_"): - self.dtype_shape = type.dtype._shape_ - 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: - self.dtype_ndim = 0 - self.dtype_shape = () - self.jax_scalar_type = wp.dtype_to_jax(type.dtype) - self.jax_ndim = type.ndim - 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: - self.dtype_ndim = 0 - self.dtype_shape = () - self.jax_scalar_type = wp.dtype_to_jax(type_to_warp(type)) - self.jax_ndim = 0 - self.warp_ndim = 0 - else: - raise TypeError(f"Invalid type for argument '{name}', expected array or scalar, got {type}") - - -class FfiLaunchDesc: - def __init__(self, static_inputs, launch_dims): - self.static_inputs = static_inputs - self.launch_dims = launch_dims - - -class FfiKernel: - def __init__(self, kernel, num_outputs, vmap_method, launch_dims, output_dims, in_out_argnames): - 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.first_array_arg = None - self.launch_id = 0 - self.launch_descriptors = {} - - in_out_argnames_list = in_out_argnames or [] - in_out_argnames = set(in_out_argnames_list) - if len(in_out_argnames_list) != len(in_out_argnames): - raise AssertionError("in_out_argnames must not contain duplicate names") - - self.num_kernel_args = len(kernel.adj.args) - self.num_in_out = len(in_out_argnames) - self.num_inputs = self.num_kernel_args - num_outputs + self.num_in_out - if self.num_outputs < 1: - raise ValueError("At least one output is required") - if self.num_outputs > self.num_kernel_args: - raise ValueError("Number of outputs cannot be greater than the number of kernel arguments") - if self.num_outputs < self.num_in_out: - raise ValueError("Number of outputs cannot be smaller than the number of in_out_argnames") - - # process input args - self.input_args = [] - for i in range(self.num_inputs): - arg_name = kernel.adj.args[i].label - arg = FfiArg(arg_name, kernel.adj.args[i].type, arg_name in in_out_argnames) - if arg_name in in_out_argnames: - in_out_argnames.remove(arg_name) - if arg.is_array: - # keep track of the first input array argument - if self.first_array_arg is None: - self.first_array_arg = i - self.input_args.append(arg) - - # process output args - self.output_args = [] - for i in range(self.num_inputs, self.num_kernel_args): - arg_name = kernel.adj.args[i].label - if arg_name in in_out_argnames: - raise AssertionError( - f"Expected an output-only argument for argument {arg_name}." - " in_out arguments should be placed before output-only arguments." - ) - arg = FfiArg(arg_name, kernel.adj.args[i].type, False) - if not arg.is_array: - raise TypeError("All output arguments must be arrays") - self.output_args.append(arg) - - if in_out_argnames: - raise ValueError(f"in_out_argnames: '{in_out_argnames}' did not match any function argument names.") - - # Build input output aliases. - out_id = 0 - input_output_aliases = {} - for in_id, arg in enumerate(self.input_args): - if not arg.in_out: - continue - input_output_aliases[in_id] = out_id - out_id += 1 - self.input_output_aliases = input_output_aliases - - # register the callback - FFI_CCALLFUNC = ctypes.CFUNCTYPE(ctypes.c_void_p, ctypes.POINTER(XLA_FFI_CallFrame)) - self.callback_func = FFI_CCALLFUNC(lambda call_frame: self.ffi_callback(call_frame)) - ffi_ccall_address = ctypes.cast(self.callback_func, ctypes.c_void_p) - ffi_capsule = jax.ffi.pycapsule(ffi_ccall_address.value) - jax.ffi.register_ffi_target(self.name, ffi_capsule, platform="CUDA") - - def __call__(self, *args, output_dims=None, launch_dims=None, vmap_method=None): - num_inputs = len(args) - if num_inputs != self.num_inputs: - raise ValueError(f"Expected {self.num_inputs} inputs, but got {num_inputs}") - - # default argument fallback - if launch_dims is None: - launch_dims = self.launch_dims - if output_dims is None: - output_dims = self.output_dims - if vmap_method is None: - vmap_method = self.vmap_method - - # output types - out_types = [] - - # process inputs - static_inputs = {} - for i in range(num_inputs): - input_arg = self.input_args[i] - input_value = args[i] - if input_arg.is_array: - # check dtype - if input_value.dtype != input_arg.jax_scalar_type: - raise TypeError( - f"Invalid data type for array argument '{input_arg.name}', expected {input_arg.jax_scalar_type}, got {input_value.dtype}" - ) - # check ndim - if input_value.ndim != input_arg.jax_ndim: - raise TypeError( - f"Invalid dimensionality for array argument '{input_arg.name}', expected {input_arg.jax_ndim} dimensions, got {input_value.ndim}" - ) - # check inner dims - for d in range(input_arg.dtype_ndim): - if input_value.shape[input_arg.type.ndim + d] != input_arg.dtype_shape[d]: - raise TypeError( - f"Invalid inner dimensions for array argument '{input_arg.name}', expected {input_arg.dtype_shape}, got {input_value.shape[-input_arg.dtype_ndim :]}" - ) - else: - # make sure scalar is not a traced variable, should be static - if isinstance(input_value, jax.core.Tracer): - raise ValueError(f"Argument '{input_arg.name}' must be a static value") - # stash the value to be retrieved by callback - static_inputs[input_arg.name] = input_arg.type(input_value) - - # append in-out arg to output types - if input_arg.in_out: - out_types.append(get_jax_output_type(input_arg, input_value.shape)) - - # launch dimensions - if launch_dims is None: - # use the shape of the first input array - if self.first_array_arg is not None: - launch_dims = get_warp_shape(self.input_args[self.first_array_arg], args[self.first_array_arg].shape) - else: - raise RuntimeError("Failed to determine launch dimensions") - elif isinstance(launch_dims, int): - launch_dims = (launch_dims,) - else: - launch_dims = tuple(launch_dims) - - # output shapes - if isinstance(output_dims, dict): - # assume a dictionary of shapes keyed on argument name - for output_arg in self.output_args: - dims = output_dims.get(output_arg.name) - if dims is None: - raise ValueError(f"Missing output dimensions for argument '{output_arg.name}'") - out_types.append(get_jax_output_type(output_arg, dims)) - else: - if output_dims is None: - # use launch dimensions - output_dims = launch_dims - elif isinstance(output_dims, int): - output_dims = (output_dims,) - # assume same dimensions for all outputs - for output_arg in self.output_args: - out_types.append(get_jax_output_type(output_arg, output_dims)) - - call = jax.ffi.ffi_call( - self.name, - out_types, - vmap_method=vmap_method, - 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) - - # save launch data to be retrieved by callback - launch_id = self.launch_id - self.launch_descriptors[launch_id] = FfiLaunchDesc(static_inputs, launch_dims) - self.launch_id += 1 - - return call(*args, launch_id=launch_id) - - def ffi_callback(self, call_frame): - 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 - # if the metadata extension is present. - extension = call_frame.contents.extension_start - if extension: - # Try to set the version metadata. - if extension.contents.type == XLA_FFI_Extension_Type.Metadata: - 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. - metadata_ext.contents.metadata.contents.traits = ( - XLA_FFI_Handler_TraitsBits.COMMAND_BUFFER_COMPATIBLE - ) - 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] - - 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))) - - assert num_inputs == self.num_inputs - assert num_outputs == self.num_outputs - - 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) - - 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) - 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.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" - - # 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, - ) - - except Exception as e: - print(traceback.format_exc()) - return create_ffi_error( - call_frame.contents.api, XLA_FFI_Error_Code.UNKNOWN, f"FFI callback error: {type(e).__name__}: {e}" - ) - - -class FfiCallDesc: - def __init__(self, static_inputs): - self.static_inputs = static_inputs - - -class FfiCallable: - def __init__(self, func, num_outputs, graph_mode, vmap_method, output_dims, in_out_argnames, graph_cache_max): - 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.first_array_arg = None - self.call_id = 0 - self.call_descriptors = {} - - # LRU cache of graphs captured by Warp - self._graph_cache_max = graph_cache_max - self.captures = collections.OrderedDict() - - in_out_argnames_list = in_out_argnames or [] - in_out_argnames = set(in_out_argnames_list) - if len(in_out_argnames_list) != len(in_out_argnames): - raise AssertionError("in_out_argnames must not contain duplicate names") - - # get arguments and annotations - argspec = get_full_arg_spec(func) - - num_args = len(argspec.args) - self.num_in_out = len(in_out_argnames) - self.num_inputs = num_args - num_outputs + self.num_in_out - if self.num_outputs < 1: - raise ValueError("At least one output is required") - if self.num_outputs > num_args: - raise ValueError("Number of outputs cannot be greater than the number of kernel arguments") - if self.num_outputs < self.num_in_out: - raise ValueError("Number of outputs cannot be smaller than the number of in_out_argnames") - - if len(argspec.annotations) < num_args: - raise RuntimeError(f"Incomplete argument annotations on function {self.name}") - - # parse type annotations - self.args = [] - arg_idx = 0 - for arg_name, arg_type in argspec.annotations.items(): - if arg_name == "return": - if arg_type is not None: - raise TypeError("Function must not return a value") - continue - else: - arg = FfiArg(arg_name, arg_type, arg_name in in_out_argnames) - if arg_name in in_out_argnames: - in_out_argnames.remove(arg_name) - if arg.is_array: - if arg_idx < self.num_inputs and self.first_array_arg is None: - self.first_array_arg = arg_idx - self.args.append(arg) - - if arg.in_out and arg_idx >= self.num_inputs: - raise AssertionError( - f"Expected an output-only argument for argument {arg_name}." - " in_out arguments should be placed before output-only arguments." - ) - - arg_idx += 1 - - if in_out_argnames: - raise ValueError(f"in_out_argnames: '{in_out_argnames}' did not match any function argument names.") - - self.input_args = self.args[: self.num_inputs] # includes in-out args - self.output_args = self.args[self.num_inputs :] # pure output args - - # Buffer indices for array arguments in callback. - # In-out buffers are the same pointers in the XLA call frame, - # so we only include them for inputs and skip them for outputs. - self.array_input_indices = [i for i, arg in enumerate(self.input_args) if arg.is_array] - self.array_output_indices = list(range(self.num_in_out, self.num_outputs)) - - # Build input output aliases. - out_id = 0 - input_output_aliases = {} - for in_id, arg in enumerate(self.input_args): - if not arg.in_out: - continue - input_output_aliases[in_id] = out_id - out_id += 1 - self.input_output_aliases = input_output_aliases - - # register the callback - FFI_CCALLFUNC = ctypes.CFUNCTYPE(ctypes.c_void_p, ctypes.POINTER(XLA_FFI_CallFrame)) - self.callback_func = FFI_CCALLFUNC(lambda call_frame: self.ffi_callback(call_frame)) - ffi_ccall_address = ctypes.cast(self.callback_func, ctypes.c_void_p) - ffi_capsule = jax.ffi.pycapsule(ffi_ccall_address.value) - jax.ffi.register_ffi_target(self.name, ffi_capsule, platform="CUDA") - - def __call__(self, *args, output_dims=None, vmap_method=None): - num_inputs = len(args) - if num_inputs != self.num_inputs: - input_names = ", ".join(arg.name for arg in self.input_args) - s = "" if self.num_inputs == 1 else "s" - raise ValueError(f"Expected {self.num_inputs} input{s} ({input_names}), but got {num_inputs}") - - # default argument fallback - if vmap_method is None: - vmap_method = self.vmap_method - if output_dims is None: - output_dims = self.output_dims - - # output types - out_types = [] - - # process inputs - static_inputs = {} - for i in range(num_inputs): - input_arg = self.input_args[i] - input_value = args[i] - if input_arg.is_array: - # check dtype - if input_value.dtype != input_arg.jax_scalar_type: - raise TypeError( - f"Invalid data type for array argument '{input_arg.name}', expected {input_arg.jax_scalar_type}, got {input_value.dtype}" - ) - # check ndim - if input_value.ndim != input_arg.jax_ndim: - raise TypeError( - f"Invalid dimensionality for array argument '{input_arg.name}', expected {input_arg.jax_ndim} dimensions, got {input_value.ndim}" - ) - # check inner dims - for d in range(input_arg.dtype_ndim): - if input_value.shape[input_arg.type.ndim + d] != input_arg.dtype_shape[d]: - raise TypeError( - f"Invalid inner dimensions for array argument '{input_arg.name}', expected {input_arg.dtype_shape}, got {input_value.shape[-input_arg.dtype_ndim :]}" - ) - else: - # make sure scalar is not a traced variable, should be static - if isinstance(input_value, jax.core.Tracer): - raise ValueError(f"Argument '{input_arg.name}' must be a static value") - # stash the value to be retrieved by callback - static_inputs[input_arg.name] = input_arg.type(input_value) - - # append in-out arg to output types - if input_arg.in_out: - out_types.append(get_jax_output_type(input_arg, input_value.shape)) - - # output shapes - if isinstance(output_dims, dict): - # assume a dictionary of shapes keyed on argument name - for output_arg in self.output_args: - dims = output_dims.get(output_arg.name) - if dims is None: - raise ValueError(f"Missing output dimensions for argument '{output_arg.name}'") - out_types.append(get_jax_output_type(output_arg, dims)) - else: - if output_dims is None: - if self.first_array_arg is None: - raise ValueError("Unable to determine output dimensions") - output_dims = get_warp_shape(self.input_args[self.first_array_arg], args[self.first_array_arg].shape) - elif isinstance(output_dims, int): - output_dims = (output_dims,) - # assume same dimensions for all outputs - for output_arg in self.output_args: - out_types.append(get_jax_output_type(output_arg, output_dims)) - - call = jax.ffi.ffi_call( - self.name, - out_types, - vmap_method=vmap_method, - input_output_aliases=self.input_output_aliases, - # has_side_effect=True, # force this function to execute even if outputs aren't used - ) - - # load the module - # 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) - - # save call data to be retrieved by callback - call_id = self.call_id - self.call_descriptors[call_id] = FfiCallDesc(static_inputs) - self.call_id += 1 - return call(*args, call_id=call_id) - - def ffi_callback(self, call_frame): - 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 - # if the metadata extension is present. - extension = call_frame.contents.extension_start - if extension: - # Try to set the version metadata. - if extension.contents.type == XLA_FFI_Extension_Type.Metadata: - 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 self.graph_mode is GraphMode.JAX: - metadata_ext.contents.metadata.contents.traits = ( - XLA_FFI_Handler_TraitsBits.COMMAND_BUFFER_COMPATIBLE - ) - 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] - - 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))) - - assert num_inputs == self.num_inputs - assert num_outputs == self.num_outputs - - 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) - - # 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 - - 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()}") - - # update the graph cache to keep recently used graphs alive - self.captures.move_to_end(capture_key) - - # early out - return - - device = wp.device_from_jax(get_jax_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 = 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: - 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()) - return create_ffi_error( - call_frame.contents.api, XLA_FFI_Error_Code.UNKNOWN, f"FFI callback error: {type(e).__name__}: {e}" - ) - - return None - - @property - def graph_cache_max(self) -> int | None: - return self._graph_cache_max - - @graph_cache_max.setter - def graph_cache_max(self, value: int | None): - if value != self._graph_cache_max: - if value is not None and (self._graph_cache_max is None or value < self._graph_cache_max): - # trim the cache if needed - while len(self.captures) > value: - self.captures.popitem(last=False) - self._graph_cache_max = value - - @property - def graph_cache_size(self) -> int: - 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 -): - """Create a JAX callback from a Warp kernel. - - NOTE: This is an experimental feature under development. - - Args: - kernel: The Warp kernel to launch. - num_outputs: Specify the number of output arguments if greater than 1. - This must include the number of ``in_out_arguments``. - vmap_method: String specifying how the callback transforms under ``vmap()``. - This argument can also be specified for individual calls. - launch_dims: Specify the default kernel launch dimensions. If None, launch - dimensions are inferred from the shape of the first array argument. - This argument can also be specified for individual calls. - 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. - - 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. - """ - - 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, - ) - - 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 - - return _FFI_KERNEL_REGISTRY[key] - - -def jax_callable( - func: Callable, - num_outputs: int = 1, - graph_compatible: Optional[bool] = None, # deprecated - graph_mode: GraphMode = GraphMode.JAX, - vmap_method: Optional[str] = "broadcast_all", - output_dims=None, - in_out_argnames=None, - graph_cache_max: int | None = None, -): - """Create a JAX callback from an annotated Python function. - - The Python function arguments must have type annotations like Warp kernels. - - NOTE: This is an experimental feature under development. - - Args: - func: The Python function to call. - num_outputs: Specify the number of output arguments if greater than 1. - This must include the number of ``in_out_arguments``. - graph_compatible: Whether the function can be called during CUDA graph capture. - This argument is deprecated, use ``graph_mode`` instead. - graph_mode: CUDA graph capture mode. - ``GraphMode.JAX`` (default): Let JAX capture the graph, which may be used as a subgraph in an enclosing JAX capture. - ``GraphMode.WARP``: Let Warp capture the graph. Use this mode when the callable cannot be used as a subgraph, - such as when the callable uses conditional graph nodes. - ``GraphMode.NONE``: Disable graph capture. Use when the callable performs operations that are not legal in a graph, - such as host synchronization. - 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. - 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. - 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``. - - 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 ``func`` definition. - - There must be at least one output or input-output argument. - - Only the CUDA backend is supported. - """ - - check_jax_version() - - if graph_compatible is not None: - wp.utils.warn( - "The `graph_compatible` argument is deprecated, use `graph_mode` instead.", - DeprecationWarning, - stacklevel=3, - ) - if graph_compatible is False: - graph_mode = GraphMode.NONE - - if graph_cache_max is None: - graph_cache_max = jax_callable_default_graph_cache_max - - # Note: we don't include graph_cache_max in the key, it is applied below. - key = ( - func, - num_outputs, - graph_mode, - vmap_method, - tuple(sorted(output_dims.items())) if output_dims else output_dims, - ) - - with _FFI_REGISTRY_LOCK: - callable = _FFI_CALLABLE_REGISTRY.get(key) - if callable is None: - callable = FfiCallable( - func, - num_outputs, - graph_mode, - vmap_method, - output_dims, - in_out_argnames, - graph_cache_max, - ) - _FFI_CALLABLE_REGISTRY[key] = callable - else: - # make sure we're using the latest graph cache max - callable.graph_cache_max = graph_cache_max - - return callable - - -def clear_jax_callable_graph_cache(callable: FfiCallable | None = None): - """Clear the graph cache of the given callable or all callables if ``None``.""" - - if callable is not None: - callable.captures.clear() - else: - # apply to all callables - with _FFI_REGISTRY_LOCK: - for callable in _FFI_CALLABLE_REGISTRY.values(): - callable.captures.clear() - - -############################################################################### -# -# Generic FFI callbacks for Python functions of the form -# func(inputs, outputs, attrs, ctx) -# -############################################################################### - - -def register_ffi_callback(name: str, func: Callable, graph_compatible: bool = True) -> None: - """Create a JAX callback from a Python function. - - The Python function must have the form ``func(inputs, outputs, attrs, ctx)``. - - NOTE: This is an experimental feature under development. - - Args: - name: A unique FFI callback name. - func: The Python function to call. - graph_compatible: Whether the function can be called during CUDA graph capture. - """ - - check_jax_version() - - # TODO check that the name is not already registered - - def ffi_callback(call_frame): - try: - extension = call_frame.contents.extension_start - # On the first call, XLA runtime will query the API version and traits - # metadata using the |extension| field. Let us respond to that query - # if the metadata extension is present. - if extension: - # Try to set the version metadata. - if extension.contents.type == XLA_FFI_Extension_Type.Metadata: - 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: - # Turn on CUDA graphs for this handler. - metadata_ext.contents.metadata.contents.traits = ( - XLA_FFI_Handler_TraitsBits.COMMAND_BUFFER_COMPATIBLE - ) - return None - - 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)] - - 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) - - func(inputs, outputs, attrs, ctx) - except Exception as e: - print(traceback.format_exc()) - return create_ffi_error( - call_frame.contents.api, XLA_FFI_Error_Code.UNKNOWN, f"FFI callback error: {type(e).__name__}: {e}" - ) - - return None - - FFI_CCALLFUNC = ctypes.CFUNCTYPE(ctypes.c_void_p, ctypes.POINTER(XLA_FFI_CallFrame)) - callback_func = FFI_CCALLFUNC(ffi_callback) - with _FFI_REGISTRY_LOCK: - _FFI_CALLBACK_REGISTRY[name] = callback_func - ffi_ccall_address = ctypes.cast(callback_func, ctypes.c_void_p) - ffi_capsule = jax.ffi.pycapsule(ffi_ccall_address.value) - jax.ffi.register_ffi_target(name, ffi_capsule, platform="CUDA") - - -############################################################################### -# -# Utilities -# -############################################################################### - -# ensure unique FFI callback names -ffi_name_counts = {} - - -def generate_unique_name(func) -> str: - key = make_full_qualified_name(func) - unique_id = ffi_name_counts.get(key, 0) - ffi_name_counts[key] = unique_id + 1 - return f"{key}_{unique_id}" - - -def get_warp_shape(arg, dims): - if arg.dtype_ndim > 0: - # vector/matrix array - return dims[: arg.warp_ndim] - else: - # scalar array - return dims - - -def get_jax_output_type(arg, dims): - if isinstance(dims, int): - dims = (dims,) - - ndim = len(dims) - - if arg.dtype_ndim > 0: - # vector/matrix array - if ndim == arg.warp_ndim: - return jax.ShapeDtypeStruct((*dims, *arg.dtype_shape), arg.jax_scalar_type) - elif ndim == arg.jax_ndim: - # make sure inner dimensions match - inner_dims = dims[-arg.dtype_ndim :] - for i in range(arg.dtype_ndim): - if inner_dims[i] != arg.dtype_shape[i]: - raise ValueError(f"Invalid output dimensions for argument '{arg.name}': {dims}") - return jax.ShapeDtypeStruct(dims, arg.jax_scalar_type) - else: - raise ValueError(f"Invalid output dimensions for argument '{arg.name}': {dims}") - else: - # scalar array - if ndim != arg.warp_ndim: - raise ValueError(f"Invalid output dimensions for argument '{arg.name}': {dims}") - return jax.ShapeDtypeStruct(dims, arg.jax_scalar_type) + return get_deprecated_api(_ffi, "wp.jax_experimental", name) diff --git a/mjx/mujoco/mjx/third_party/warp/jax_experimental/xla_ffi.py b/mjx/mujoco/mjx/third_party/warp/jax_experimental/xla_ffi.py index 8e311143..93f57b3c 100644 --- a/mjx/mujoco/mjx/third_party/warp/jax_experimental/xla_ffi.py +++ b/mjx/mujoco/mjx/third_party/warp/jax_experimental/xla_ffi.py @@ -13,612 +13,12 @@ # See the License for the specific language governing permissions and # limitations under the License. -import ctypes -import enum +# TODO: Remove after cleaning up the public API. -import jax.numpy as jnp -import numpy as np +from warp._src.jax_experimental import xla_ffi as _xla_ffi -import warp as wp -####################################################################### -# ctypes structures and enums for XLA's FFI API: -# https://github.com/openxla/xla/blob/a1a5e62fbffa3a3b6c409d72607456cf5b353a22/xla/ffi/api/c_api.h -####################################################################### +def __getattr__(name): + from warp._src.utils import get_deprecated_api - -# typedef enum { -# XLA_FFI_Extension_Metadata = 1, -# } XLA_FFI_Extension_Type; -class XLA_FFI_Extension_Type(enum.IntEnum): - Metadata = 1 - - -# typedef struct XLA_FFI_Extension_Base { -# size_t struct_size; -# XLA_FFI_Extension_Type type; -# struct XLA_FFI_Extension_Base* next; -# } XLA_FFI_Extension_Base; -class XLA_FFI_Extension_Base(ctypes.Structure): - pass - - -XLA_FFI_Extension_Base._fields_ = [ - ("struct_size", ctypes.c_size_t), - ("type", ctypes.c_int), # XLA_FFI_Extension_Type - ("next", ctypes.POINTER(XLA_FFI_Extension_Base)), -] - - -# typedef enum { -# XLA_FFI_ExecutionStage_INSTANTIATE = 0, -# XLA_FFI_ExecutionStage_PREPARE = 1, -# XLA_FFI_ExecutionStage_INITIALIZE = 2, -# XLA_FFI_ExecutionStage_EXECUTE = 3, -# } XLA_FFI_ExecutionStage; -class XLA_FFI_ExecutionStage(enum.IntEnum): - INSTANTIATE = 0 - PREPARE = 1 - INITIALIZE = 2 - EXECUTE = 3 - - -# typedef enum { -# XLA_FFI_DataType_INVALID = 0, -# XLA_FFI_DataType_PRED = 1, -# XLA_FFI_DataType_S8 = 2, -# XLA_FFI_DataType_S16 = 3, -# XLA_FFI_DataType_S32 = 4, -# XLA_FFI_DataType_S64 = 5, -# XLA_FFI_DataType_U8 = 6, -# XLA_FFI_DataType_U16 = 7, -# XLA_FFI_DataType_U32 = 8, -# XLA_FFI_DataType_U64 = 9, -# XLA_FFI_DataType_F16 = 10, -# XLA_FFI_DataType_F32 = 11, -# XLA_FFI_DataType_F64 = 12, -# XLA_FFI_DataType_BF16 = 16, -# XLA_FFI_DataType_C64 = 15, -# XLA_FFI_DataType_C128 = 18, -# XLA_FFI_DataType_TOKEN = 17, -# XLA_FFI_DataType_F8E5M2 = 19, -# XLA_FFI_DataType_F8E3M4 = 29, -# XLA_FFI_DataType_F8E4M3 = 28, -# XLA_FFI_DataType_F8E4M3FN = 20, -# XLA_FFI_DataType_F8E4M3B11FNUZ = 23, -# XLA_FFI_DataType_F8E5M2FNUZ = 24, -# XLA_FFI_DataType_F8E4M3FNUZ = 25, -# XLA_FFI_DataType_F4E2M1FN = 32, -# XLA_FFI_DataType_F8E8M0FNU = 33, -# } XLA_FFI_DataType; -class XLA_FFI_DataType(enum.IntEnum): - INVALID = 0 - PRED = 1 - S8 = 2 - S16 = 3 - S32 = 4 - S64 = 5 - U8 = 6 - U16 = 7 - U32 = 8 - U64 = 9 - F16 = 10 - F32 = 11 - F64 = 12 - BF16 = 16 - C64 = 15 - C128 = 18 - TOKEN = 17 - F8E5M2 = 19 - F8E3M4 = 29 - F8E4M3 = 28 - F8E4M3FN = 20 - F8E4M3B11FNUZ = 23 - F8E5M2FNUZ = 24 - F8E4M3FNUZ = 25 - F4E2M1FN = 32 - F8E8M0FNU = 33 - - -# struct XLA_FFI_Buffer { -# size_t struct_size; -# XLA_FFI_Extension_Base* extension_start; -# -# XLA_FFI_DataType dtype; -# void* data; -# int64_t rank; -# int64_t* dims; // length == rank -# }; -class XLA_FFI_Buffer(ctypes.Structure): - _fields_ = ( - ("struct_size", ctypes.c_size_t), - ("extension_start", ctypes.POINTER(XLA_FFI_Extension_Base)), - ("dtype", ctypes.c_int), # XLA_FFI_DataType - ("data", ctypes.c_void_p), - ("rank", ctypes.c_int64), - ("dims", ctypes.POINTER(ctypes.c_int64)), - ) - - -# typedef enum { -# XLA_FFI_ArgType_BUFFER = 1, -# } XLA_FFI_ArgType; -class XLA_FFI_ArgType(enum.IntEnum): - BUFFER = 1 - - -# typedef enum { -# XLA_FFI_RetType_BUFFER = 1, -# } XLA_FFI_RetType; -class XLA_FFI_RetType(enum.IntEnum): - BUFFER = 1 - - -# struct XLA_FFI_Args { -# size_t struct_size; -# XLA_FFI_Extension_Base* extension_start; -# int64_t size; -# XLA_FFI_ArgType* types; // length == size -# void** args; // length == size -# }; -class XLA_FFI_Args(ctypes.Structure): - _fields_ = ( - ("struct_size", ctypes.c_size_t), - ("extension_start", ctypes.POINTER(XLA_FFI_Extension_Base)), - ("size", ctypes.c_int64), - ("types", ctypes.POINTER(ctypes.c_int)), # XLA_FFI_ArgType* - ("args", ctypes.POINTER(ctypes.c_void_p)), - ) - - -# struct XLA_FFI_Rets { -# size_t struct_size; -# XLA_FFI_Extension_Base* extension_start; -# int64_t size; -# XLA_FFI_RetType* types; // length == size -# void** rets; // length == size -# }; -class XLA_FFI_Rets(ctypes.Structure): - _fields_ = ( - ("struct_size", ctypes.c_size_t), - ("extension_start", ctypes.POINTER(XLA_FFI_Extension_Base)), - ("size", ctypes.c_int64), - ("types", ctypes.POINTER(ctypes.c_int)), # XLA_FFI_RetType* - ("rets", ctypes.POINTER(ctypes.c_void_p)), - ) - - -# typedef struct XLA_FFI_ByteSpan { -# const char* ptr; -# size_t len; -# } XLA_FFI_ByteSpan; -class XLA_FFI_ByteSpan(ctypes.Structure): - _fields_ = ( - ("ptr", ctypes.POINTER(ctypes.c_char)), - ("len", ctypes.c_size_t), - ) - - -# typedef struct XLA_FFI_Scalar { -# XLA_FFI_DataType dtype; -# void* value; -# } XLA_FFI_Scalar; -class XLA_FFI_Scalar(ctypes.Structure): - _fields_ = ( - ("dtype", ctypes.c_int), - ("value", ctypes.c_void_p), - ) - - -# typedef struct XLA_FFI_Array { -# XLA_FFI_DataType dtype; -# size_t size; -# void* data; -# } XLA_FFI_Array; -class XLA_FFI_Array(ctypes.Structure): - _fields_ = ( - ("dtype", ctypes.c_int), - ("size", ctypes.c_size_t), - ("data", ctypes.c_void_p), - ) - - -# typedef enum { -# XLA_FFI_AttrType_ARRAY = 1, -# XLA_FFI_AttrType_DICTIONARY = 2, -# XLA_FFI_AttrType_SCALAR = 3, -# XLA_FFI_AttrType_STRING = 4, -# } XLA_FFI_AttrType; -class XLA_FFI_AttrType(enum.IntEnum): - ARRAY = 1 - DICTIONARY = 2 - SCALAR = 3 - STRING = 4 - - -# struct XLA_FFI_Attrs { -# size_t struct_size; -# XLA_FFI_Extension_Base* extension_start; -# int64_t size; -# XLA_FFI_AttrType* types; // length == size -# XLA_FFI_ByteSpan** names; // length == size -# void** attrs; // length == size -# }; -class XLA_FFI_Attrs(ctypes.Structure): - _fields_ = ( - ("struct_size", ctypes.c_size_t), - ("extension_start", ctypes.POINTER(XLA_FFI_Extension_Base)), - ("size", ctypes.c_int64), - ("types", ctypes.POINTER(ctypes.c_int)), # XLA_FFI_AttrType* - ("names", ctypes.POINTER(ctypes.POINTER(XLA_FFI_ByteSpan))), - ("attrs", ctypes.POINTER(ctypes.c_void_p)), - ) - - -# struct XLA_FFI_Api_Version { -# size_t struct_size; -# XLA_FFI_Extension_Base* extension_start; -# int major_version; // out -# int minor_version; // out -# }; -class XLA_FFI_Api_Version(ctypes.Structure): - _fields_ = ( - ("struct_size", ctypes.c_size_t), - ("extension_start", ctypes.POINTER(XLA_FFI_Extension_Base)), - ("major_version", ctypes.c_int), - ("minor_version", ctypes.c_int), - ) - - -# enum XLA_FFI_Handler_TraitsBits { -# // Calls to FFI handler are safe to trace into the command buffer. It means -# // that calls to FFI handler always launch exactly the same device operations -# // (can depend on attribute values) that can be captured and then replayed. -# XLA_FFI_HANDLER_TRAITS_COMMAND_BUFFER_COMPATIBLE = 1u << 0, -# }; -class XLA_FFI_Handler_TraitsBits(enum.IntEnum): - COMMAND_BUFFER_COMPATIBLE = 1 << 0 - - -# struct XLA_FFI_Metadata { -# size_t struct_size; -# XLA_FFI_Api_Version api_version; -# XLA_FFI_Handler_Traits traits; -# }; -class XLA_FFI_Metadata(ctypes.Structure): - _fields_ = ( - ("struct_size", ctypes.c_size_t), - ("api_version", XLA_FFI_Api_Version), # XLA_FFI_Extension_Type - ("traits", ctypes.c_uint32), # XLA_FFI_Handler_Traits - ) - - -# struct XLA_FFI_Metadata_Extension { -# XLA_FFI_Extension_Base extension_base; -# XLA_FFI_Metadata* metadata; -# }; -class XLA_FFI_Metadata_Extension(ctypes.Structure): - _fields_ = ( - ("extension_base", XLA_FFI_Extension_Base), - ("metadata", ctypes.POINTER(XLA_FFI_Metadata)), - ) - - -# typedef enum { -# XLA_FFI_Error_Code_OK = 0, -# XLA_FFI_Error_Code_CANCELLED = 1, -# XLA_FFI_Error_Code_UNKNOWN = 2, -# XLA_FFI_Error_Code_INVALID_ARGUMENT = 3, -# XLA_FFI_Error_Code_DEADLINE_EXCEEDED = 4, -# XLA_FFI_Error_Code_NOT_FOUND = 5, -# XLA_FFI_Error_Code_ALREADY_EXISTS = 6, -# XLA_FFI_Error_Code_PERMISSION_DENIED = 7, -# XLA_FFI_Error_Code_RESOURCE_EXHAUSTED = 8, -# XLA_FFI_Error_Code_FAILED_PRECONDITION = 9, -# XLA_FFI_Error_Code_ABORTED = 10, -# XLA_FFI_Error_Code_OUT_OF_RANGE = 11, -# XLA_FFI_Error_Code_UNIMPLEMENTED = 12, -# XLA_FFI_Error_Code_INTERNAL = 13, -# XLA_FFI_Error_Code_UNAVAILABLE = 14, -# XLA_FFI_Error_Code_DATA_LOSS = 15, -# XLA_FFI_Error_Code_UNAUTHENTICATED = 16 -# } XLA_FFI_Error_Code; -class XLA_FFI_Error_Code(enum.IntEnum): - OK = 0 - CANCELLED = 1 - UNKNOWN = 2 - INVALID_ARGUMENT = 3 - DEADLINE_EXCEEDED = 4 - NOT_FOUND = 5 - ALREADY_EXISTS = 6 - PERMISSION_DENIED = 7 - RESOURCE_EXHAUSTED = 8 - FAILED_PRECONDITION = 9 - ABORTED = 10 - OUT_OF_RANGE = 11 - UNIMPLEMENTED = 12 - INTERNAL = 13 - UNAVAILABLE = 14 - DATA_LOSS = 15 - UNAUTHENTICATED = 16 - - -# struct XLA_FFI_Error_Create_Args { -# size_t struct_size; -# XLA_FFI_Extension_Base* extension_start; -# const char* message; -# XLA_FFI_Error_Code errc; -# }; -class XLA_FFI_Error_Create_Args(ctypes.Structure): - _fields_ = ( - ("struct_size", ctypes.c_size_t), - ("extension_start", ctypes.POINTER(XLA_FFI_Extension_Base)), - ("message", ctypes.c_char_p), - ("errc", ctypes.c_int), - ) # XLA_FFI_Error_Code - - -XLA_FFI_Error_Create = ctypes.CFUNCTYPE(ctypes.c_void_p, ctypes.POINTER(XLA_FFI_Error_Create_Args)) - - -# struct XLA_FFI_Stream_Get_Args { -# size_t struct_size; -# XLA_FFI_Extension_Base* extension_start; -# XLA_FFI_ExecutionContext* ctx; -# void* stream; // out -# }; -class XLA_FFI_Stream_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* - ("stream", ctypes.c_void_p), - ) # // out - - -XLA_FFI_Stream_Get = ctypes.CFUNCTYPE(ctypes.c_void_p, ctypes.POINTER(XLA_FFI_Stream_Get_Args)) - - -# struct XLA_FFI_Api { -# size_t struct_size; -# XLA_FFI_Extension_Base* extension_start; -# -# XLA_FFI_Api_Version api_version; -# XLA_FFI_InternalApi* internal_api; -# -# _XLA_FFI_API_STRUCT_FIELD(XLA_FFI_Error_Create); -# _XLA_FFI_API_STRUCT_FIELD(XLA_FFI_Error_GetMessage); -# _XLA_FFI_API_STRUCT_FIELD(XLA_FFI_Error_Destroy); -# _XLA_FFI_API_STRUCT_FIELD(XLA_FFI_Handler_Register); -# _XLA_FFI_API_STRUCT_FIELD(XLA_FFI_Stream_Get); -# _XLA_FFI_API_STRUCT_FIELD(XLA_FFI_TypeId_Register); -# _XLA_FFI_API_STRUCT_FIELD(XLA_FFI_ExecutionContext_Get); -# _XLA_FFI_API_STRUCT_FIELD(XLA_FFI_State_Set); -# _XLA_FFI_API_STRUCT_FIELD(XLA_FFI_State_Get); -# _XLA_FFI_API_STRUCT_FIELD(XLA_FFI_DeviceMemory_Allocate); -# _XLA_FFI_API_STRUCT_FIELD(XLA_FFI_DeviceMemory_Free); -# _XLA_FFI_API_STRUCT_FIELD(XLA_FFI_ThreadPool_Schedule); -# _XLA_FFI_API_STRUCT_FIELD(XLA_FFI_ThreadPool_NumThreads); -# _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); -# }; -class XLA_FFI_Api(ctypes.Structure): - _fields_ = ( - ("struct_size", ctypes.c_size_t), - ("extension_start", ctypes.POINTER(XLA_FFI_Extension_Base)), - ("api_version", XLA_FFI_Api_Version), - ("internal_api", ctypes.c_void_p), # XLA_FFI_InternalApi* - ("XLA_FFI_Error_Create", XLA_FFI_Error_Create), # XLA_FFI_Error_Create - ("XLA_FFI_Error_GetMessage", ctypes.c_void_p), # XLA_FFI_Error_GetMessage - ("XLA_FFI_Error_Destroy", ctypes.c_void_p), # XLA_FFI_Error_Destroy - ("XLA_FFI_Handler_Register", ctypes.c_void_p), # XLA_FFI_Handler_Register - ("XLA_FFI_Stream_Get", XLA_FFI_Stream_Get), # XLA_FFI_Stream_Get - ("XLA_FFI_TypeId_Register", ctypes.c_void_p), # XLA_FFI_TypeId_Register - ("XLA_FFI_ExecutionContext_Get", ctypes.c_void_p), # XLA_FFI_ExecutionContext_Get - ("XLA_FFI_State_Set", ctypes.c_void_p), # XLA_FFI_State_Set - ("XLA_FFI_State_Get", ctypes.c_void_p), # XLA_FFI_State_Get - ("XLA_FFI_DeviceMemory_Allocate", ctypes.c_void_p), # XLA_FFI_DeviceMemory_Allocate - ("XLA_FFI_DeviceMemory_Free", ctypes.c_void_p), # XLA_FFI_DeviceMemory_Free - ("XLA_FFI_ThreadPool_Schedule", ctypes.c_void_p), # XLA_FFI_ThreadPool_Schedule - ("XLA_FFI_ThreadPool_NumThreads", ctypes.c_void_p), # XLA_FFI_ThreadPool_NumThreads - ("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 - ) - - -# struct XLA_FFI_CallFrame { -# size_t struct_size; -# XLA_FFI_Extension_Base* extension_start; -# const XLA_FFI_Api* api; -# XLA_FFI_ExecutionContext* ctx; -# XLA_FFI_ExecutionStage stage; -# XLA_FFI_Args args; -# XLA_FFI_Rets rets; -# XLA_FFI_Attrs attrs; -# -# // XLA FFI handler implementation can use `future` to signal a result of -# // asynchronous computation to the XLA runtime. XLA runtime will keep all -# // arguments, results and attributes alive until `future` is completed. -# XLA_FFI_Future* future; // out -# }; -class XLA_FFI_CallFrame(ctypes.Structure): - _fields_ = ( - ("struct_size", ctypes.c_size_t), - ("extension_start", ctypes.POINTER(XLA_FFI_Extension_Base)), - ("api", ctypes.POINTER(XLA_FFI_Api)), - ("ctx", ctypes.c_void_p), # XLA_FFI_ExecutionContext* - ("stage", ctypes.c_int), # XLA_FFI_ExecutionStage - ("args", XLA_FFI_Args), - ("rets", XLA_FFI_Rets), - ("attrs", XLA_FFI_Attrs), - ("future", ctypes.c_void_p), # XLA_FFI_Future* // out - ) - - -_xla_data_type_to_constructor = { - # XLA_FFI_DataType.INVALID - XLA_FFI_DataType.PRED: jnp.bool, - XLA_FFI_DataType.S8: jnp.int8, - XLA_FFI_DataType.S16: jnp.int16, - XLA_FFI_DataType.S32: jnp.int32, - XLA_FFI_DataType.S64: jnp.int64, - XLA_FFI_DataType.U8: jnp.uint8, - XLA_FFI_DataType.U16: jnp.uint16, - XLA_FFI_DataType.U32: jnp.uint32, - XLA_FFI_DataType.U64: jnp.uint64, - XLA_FFI_DataType.F16: jnp.float16, - XLA_FFI_DataType.F32: jnp.float32, - XLA_FFI_DataType.F64: jnp.float64, - XLA_FFI_DataType.BF16: jnp.bfloat16, - XLA_FFI_DataType.C64: jnp.complex64, - XLA_FFI_DataType.C128: jnp.complex128, - # XLA_FFI_DataType.TOKEN - # XLA_FFI_DataType.F4E2M1FN: jnp.float4_e2m1fn.dtype, - # XLA_FFI_DataType.F8E8M0FNU: jnp.float8_e8m0fnu.dtype, -} - -# newer types not supported by older versions -if hasattr(jnp, "float8_e5m2"): - _xla_data_type_to_constructor[XLA_FFI_DataType.F8E5M2] = jnp.float8_e5m2 -if hasattr(jnp, "float8_e3m4"): - _xla_data_type_to_constructor[XLA_FFI_DataType.F8E3M4] = jnp.float8_e3m4 -if hasattr(jnp, "float8_e4m3"): - _xla_data_type_to_constructor[XLA_FFI_DataType.F8E4M3] = jnp.float8_e4m3 -if hasattr(jnp, "float8_e4m3fn"): - _xla_data_type_to_constructor[XLA_FFI_DataType.F8E4M3FN] = jnp.float8_e4m3fn -if hasattr(jnp, "float8_e4m3b11fnuz"): - _xla_data_type_to_constructor[XLA_FFI_DataType.F8E4M3B11FNUZ] = jnp.float8_e4m3b11fnuz -if hasattr(jnp, "float8_e5m2fnuz"): - _xla_data_type_to_constructor[XLA_FFI_DataType.F8E5M2FNUZ] = jnp.float8_e5m2fnuz -if hasattr(jnp, "float8_e4m3fnuz"): - _xla_data_type_to_constructor[XLA_FFI_DataType.F8E4M3FNUZ] = jnp.float8_e4m3fnuz - - -######################################################################## -# Helpers for translating between ctypes and python types -####################################################################### - - -def decode_bytespan(span: XLA_FFI_ByteSpan): - len = span.len - chars = ctypes.cast(span.ptr, ctypes.POINTER(ctypes.c_char * len)) - return chars.contents.value.decode("utf-8") - - -def decode_scalar(scalar: XLA_FFI_Scalar): - # TODO validate if dtype supported - dtype = jnp.dtype(_xla_data_type_to_constructor[scalar.dtype]) - bytes = ctypes.string_at(scalar.value, dtype.itemsize) - return np.frombuffer(bytes, dtype=dtype).reshape(()) - - -def decode_array(array: XLA_FFI_Array): - # TODO validate if dtype supported - dtype = jnp.dtype(_xla_data_type_to_constructor[array.dtype]) - bytes = ctypes.string_at(array.data, dtype.itemsize * array.size) - return np.frombuffer(bytes, dtype=dtype) - - -def decode_attrs(attrs: XLA_FFI_Attrs): - result = {} - for i in range(attrs.size): - attr_name = decode_bytespan(attrs.names[i].contents) - attr_type = attrs.types[i] - if attr_type == XLA_FFI_AttrType.STRING: - bytespan = ctypes.cast(attrs.attrs[i], ctypes.POINTER(XLA_FFI_ByteSpan)) - attr_value = decode_bytespan(bytespan.contents) - elif attr_type == XLA_FFI_AttrType.SCALAR: - attr_value = ctypes.cast(attrs.attrs[i], ctypes.POINTER(XLA_FFI_Scalar)) - attr_value = decode_scalar(attr_value.contents) - elif attr_type == XLA_FFI_AttrType.ARRAY: - attr_value = ctypes.cast(attrs.attrs[i], ctypes.POINTER(XLA_FFI_Array)) - attr_value = decode_array(attr_value.contents) - elif attr_type == XLA_FFI_AttrType.DICTIONARY: - attr_value = ctypes.cast(attrs.attrs[i], ctypes.POINTER(XLA_FFI_Attrs)) - attr_value = decode_attrs(attr_value.contents) - else: - raise Exception("Unexpected attr type") - result[attr_name] = attr_value - return result - - -# error-string to XLA_FFI_Error -def create_ffi_error(api, errc, message): - create_args = XLA_FFI_Error_Create_Args( - ctypes.sizeof(XLA_FFI_Error_Create_Args), - ctypes.POINTER(XLA_FFI_Extension_Base)(), - ctypes.c_char_p(message.encode("utf-8")), - errc, - ) - return api.contents.XLA_FFI_Error_Create(create_args) - - -def create_invalid_argument_ffi_error(api, message): - return create_ffi_error(api, XLA_FFI_Error_Code.INVALID_ARGUMENT, message) - - -# Extract CUDA stream from XLA_FFI_CallFrame. -def get_stream_from_callframe(call_frame): - api = call_frame.api - get_stream_args = XLA_FFI_Stream_Get_Args( - ctypes.sizeof(XLA_FFI_Stream_Get_Args), ctypes.POINTER(XLA_FFI_Extension_Base)(), call_frame.ctx, None - ) - api.contents.XLA_FFI_Stream_Get(get_stream_args) - # TODO check result - return get_stream_args.stream - - -_dtype_from_ffi = { - XLA_FFI_DataType.S8: wp.int8, - XLA_FFI_DataType.S16: wp.int16, - XLA_FFI_DataType.S32: wp.int32, - XLA_FFI_DataType.S64: wp.int64, - XLA_FFI_DataType.U8: wp.uint8, - XLA_FFI_DataType.U16: wp.uint16, - XLA_FFI_DataType.U32: wp.uint32, - XLA_FFI_DataType.U64: wp.uint64, - XLA_FFI_DataType.F16: wp.float16, - XLA_FFI_DataType.F32: wp.float32, - XLA_FFI_DataType.F64: wp.float64, -} - - -def dtype_from_ffi(ffi_dtype): - return _dtype_from_ffi.get(ffi_dtype) - - -def jax_dtype_from_ffi(ffi_dtype): - return _xla_data_type_to_constructor.get(ffi_dtype) - - -# Execution context (stream, stage) -class ExecutionContext: - stage: XLA_FFI_ExecutionStage - stream: int - - def __init__(self, callframe: XLA_FFI_CallFrame): - self.stage = XLA_FFI_ExecutionStage(callframe.stage) - self.stream = get_stream_from_callframe(callframe) - - -class FfiBuffer: - dtype: str - data: int - shape: tuple[int] - - def __init__(self, xla_buffer): - # TODO check if valid - self.dtype = jnp.dtype(_xla_data_type_to_constructor[xla_buffer.dtype]) - self.shape = tuple(xla_buffer.dims[i] for i in range(xla_buffer.rank)) - self.data = xla_buffer.data - - @property - def __cuda_array_interface__(self): - return { - "shape": self.shape, - "typestr": self.dtype.char, - "data": (self.data, False), - "version": 2, - } + return get_deprecated_api(_xla_ffi, "wp.jax_experimental", name)