Import Warp v1.14.0

PiperOrigin-RevId: 938464395
Change-Id: Iac99c53566e2e62f6e75f5e0904100919043b463
This commit is contained in:
Taylor Howell
2026-06-26 02:03:11 -07:00
committed by Copybara-Service
parent 07e292bc6b
commit 56a182e67c
19 changed files with 505 additions and 218 deletions
+5 -5
View File
@@ -20,8 +20,8 @@ jax-cuda12-pjrt==0.8.3; python_version >= '3.13' \
jax-cuda12-pjrt==0.5.3; python_version < '3.13' \
--hash=sha256:04ee111eaf5fc2692978ad4a5c84d5925e42eb05c1701849ba3a53f6515400cc \
--hash=sha256:c5378306568ba0c81b230a779dd3194c9dd10339ab6360ae80928108d37e7f75
warp-lang==1.13.0 \
--hash=sha256:4375f572301991fe0fbf0af29fc84d76cd27d531432d6df8452b25088c21ea5a \
--hash=sha256:ac2479c70ad410d58deb088c2a64168792be588efc1bec22dc51f92238e8c8c3 \
--hash=sha256:476b54f0dcf6767f23305a328660a9a74d01e371b6b7c3d77e03e18b4a1bb1a5 \
--hash=sha256:47975ea07252d45a4d09d2d1a6cffc55002fa7fbde771c8dceb1baf4b32d4fb8
warp-lang==1.14.0 \
--hash=sha256:12656050545cc77bf9b9b155399496c1a6279b5b6c59e407507d6858a2beb4a2 \
--hash=sha256:70cd127d0e9109417099649fedf9d00f39f1307ccb7a6e9fb87661337868d7de \
--hash=sha256:f482787e8da9c9ef045601fde99095e16d604fbcc3cbb4a1e0cef0769388b316 \
--hash=sha256:936b49ec78237f9760e58cbe9c46ee6f4244aefbd62071c4fa9fd3b313dfa878
+1 -1
View File
@@ -531,7 +531,7 @@ def put_model(
device: which device to use - if unspecified picks the default device
impl: implementation to use
graph_mode: CUDA graph capture mode (for Warp only). Use GraphMode enum from
warp._src.jax_experimental.ffi. GraphMode.WARP is the default mode.
warp._src.jax.ffi.GraphMode.WARP is the default mode.
keepalive_refs: optional dict to store references to underlying MuJoCo
objects, preventing them from being garbage collected. Required for CPP
impl to keep the model alive.
@@ -29,7 +29,8 @@ from mujoco.mjx.codegen import file
import mujoco.mjx.third_party.mujoco_warp as mjwarp
import numpy as np
import warp as wp
from mujoco.mjx.third_party.warp._src.jax_experimental import ffi
from warp._src.jax.ffi import FfiArg
from warp import JaxCallableGraphMode as GraphMode
_MJX_WARP_TYPES_OUT_FPATH = flags.DEFINE_string(
@@ -116,7 +117,7 @@ def _get_target_annotation_node(
if annotation in (int, float, bool):
return _ast_parse_type(annotation.__name__)
if annotation is ffi.GraphMode:
if annotation is GraphMode:
return _ast_parse_type('GraphMode')
if isinstance(annotation, type) and issubclass(annotation, enum.Enum):
@@ -273,7 +274,8 @@ if typing.TYPE_CHECKING:
pass
else:
try:
from warp._src.jax_experimental.ffi import GraphMode
from mujoco.mjx.third_party.warp._src.jax import ffi as warp_ffi
GraphMode = warp_ffi.JaxCallableGraphMode
from mujoco.mjx.third_party.mujoco_warp._src import types as mjwp_types
Callback = mjwp_types.Callback
except ImportError:
@@ -501,7 +503,7 @@ def _to_jax_ndim(name: str, wp_type: Any) -> int:
if typing.get_args(wp_type)[1] != ...:
raise NotImplementedError('Only variadic tuples are supported.')
return -1 # signals that dim should be untouched in downstream code.
ffi_arg = ffi.FfiArg(name, wp_type)
ffi_arg = FfiArg(name, wp_type)
return ffi_arg.jax_ndim
@@ -573,7 +575,7 @@ def main(argv):
write_core_cls('Statistic', target_fpath, mjx_types_fpath, set_diff=False)
write_core_cls(
'Option', target_fpath, mjx_types_fpath,
extra_annotations={'graph_mode': ffi.GraphMode},
extra_annotations={'graph_mode': GraphMode},
)
write_core_cls('Model', target_fpath, mjx_types_fpath)
write_core_cls('Data', target_fpath, mjx_types_fpath, flatten_fields=True)
+47 -27
View File
@@ -555,6 +555,26 @@ def render(m: Model, d: Data, rc: RenderContext):
cast_ray_first_hit = _make_cast_ray(geom_ray_types, first_hit=True)
compute_lighting = _make_compute_lighting(cast_ray_first_hit)
# Extract static parameters to avoid capturing rc and m in the kernel closure.
static_use_precomputed_rays = rc.use_precomputed_rays
static_znear = rc.znear
static_enable_backface_culling = rc.enable_backface_culling
static_render_skybox = rc.render_skybox
static_skybox_tex_id = rc.skybox_tex_id
static_skybox_face_width = rc.skybox_face_width
static_use_textures = rc.use_textures
static_enable_specular = rc.enable_specular
static_enable_emission = rc.enable_emission
static_use_ambient_lighting = rc.use_ambient_lighting
static_headlight_active = rc.headlight_active
static_headlight_ambient = rc.headlight_ambient
static_headlight_diffuse = rc.headlight_diffuse
static_headlight_specular = rc.headlight_specular
static_enable_per_light_ambient = rc.enable_per_light_ambient
static_light_attenuation_is_default = rc.light_attenuation_is_default
static_has_spot_lights = rc.has_spot_lights
static_nlight = m.nlight
@wp.kernel(module="unique", enable_backward=False)
def _render_megakernel(
# Model:
@@ -649,7 +669,7 @@ def render(m: Model, d: Data, rc: RenderContext):
# Map active camera index to MuJoCo camera ID
mujoco_cam_id = cam_id_map[camid]
if wp.static(rc.use_precomputed_rays):
if wp.static(static_use_precomputed_rays):
ray_dir_local_cam = ray[rayid]
else:
img_w = cam_res[camid][0]
@@ -665,7 +685,7 @@ def render(m: Model, d: Data, rc: RenderContext):
img_h,
px,
py,
wp.static(rc.znear),
wp.static(static_znear),
)
ray_origin_world = cam_xpos_in[worldid, mujoco_cam_id]
@@ -697,7 +717,7 @@ def render(m: Model, d: Data, rc: RenderContext):
ray_origin_world,
ray_dir_world,
float(MJ_MAXVAL),
wp.static(rc.enable_backface_culling),
wp.static(static_enable_backface_culling),
)
if render_seg[camid] and geom_id != -1:
@@ -708,10 +728,10 @@ def render(m: Model, d: Data, rc: RenderContext):
# Early Out
if geom_id == -1:
if wp.static(rc.render_skybox) and render_rgb[camid]:
if wp.static(static_render_skybox) and render_rgb[camid]:
skybox_color = sample_skybox(
textures[wp.static(rc.skybox_tex_id)],
wp.static(1.0 / float(rc.skybox_face_width)),
textures[wp.static(static_skybox_tex_id)],
wp.static(1.0 / float(static_skybox_face_width)),
ray_dir_world,
)
rgb_out[worldid, rgb_adr[camid] + rayid_local] = pack_rgba_to_uint32(
@@ -745,7 +765,7 @@ def render(m: Model, d: Data, rc: RenderContext):
base_color = wp.vec3(color[0], color[1], color[2])
if wp.static(rc.use_textures):
if wp.static(static_use_textures):
if geom_id != -2:
mat_id = geom_matid[worldid % geom_matid.shape[0], geom_id]
if mat_id >= 0:
@@ -773,27 +793,27 @@ def render(m: Model, d: Data, rc: RenderContext):
mat_spec = DEFAULT_MAT_SPECULAR
mat_shin_exp = DEFAULT_MAT_SHININESS_EXPONENT
mat_emis = DEFAULT_MAT_EMISSION
if wp.static(rc.enable_specular or rc.enable_emission):
if wp.static(static_enable_specular or static_enable_emission):
if geom_id != -2:
mat_id_for_spec = geom_matid[worldid % geom_matid.shape[0], geom_id]
if mat_id_for_spec >= 0:
if wp.static(rc.enable_specular):
if wp.static(static_enable_specular):
mat_spec = mat_specular[worldid % mat_specular.shape[0], mat_id_for_spec]
mat_shin_exp = mat_shininess[worldid % mat_shininess.shape[0], mat_id_for_spec] * MAX_SHININESS
if wp.static(rc.enable_emission):
if wp.static(static_enable_emission):
mat_emis = mat_emission[worldid % mat_emission.shape[0], mat_id_for_spec]
result = wp.vec3(0.0)
if wp.static(rc.enable_emission):
if wp.static(static_enable_emission):
result = base_color * mat_emis
if wp.static(rc.use_ambient_lighting):
if wp.static(rc.headlight_active):
result = result + wp.cw_mul(base_color, wp.static(rc.headlight_ambient))
elif wp.static(m.nlight == 0):
if wp.static(static_use_ambient_lighting):
if wp.static(static_headlight_active):
result = result + wp.cw_mul(base_color, wp.static(static_headlight_ambient))
elif wp.static(static_nlight == 0):
result = result + base_color * NO_LIGHT_AMBIENT_FALLBACK
if wp.static(rc.enable_per_light_ambient):
for light_index in range(wp.static(m.nlight)):
if wp.static(static_enable_per_light_ambient):
for light_index in range(wp.static(static_nlight)):
if light_active[worldid % light_active.shape[0], light_index]:
result = result + wp.cw_mul(base_color, light_ambient[worldid % light_ambient.shape[0], light_index])
@@ -810,7 +830,7 @@ def render(m: Model, d: Data, rc: RenderContext):
light_diffuse_worldid = light_diffuse[worldid % light_diffuse.shape[0]]
light_specular_worldid = light_specular[worldid % light_specular.shape[0]]
# Apply Lighting for each light
for light_index in range(wp.static(m.nlight)):
for light_index in range(wp.static(static_nlight)):
diff_rgb, spec_rgb = compute_lighting(
geom_type,
geom_dataid,
@@ -849,15 +869,15 @@ def render(m: Model, d: Data, rc: RenderContext):
view_dir,
mat_spec,
mat_shin_exp,
wp.static(rc.enable_backface_culling),
wp.static(rc.enable_specular),
wp.static(rc.light_attenuation_is_default),
wp.static(rc.has_spot_lights),
wp.static(static_enable_backface_culling),
wp.static(static_enable_specular),
wp.static(static_light_attenuation_is_default),
wp.static(static_has_spot_lights),
)
result = result + wp.cw_mul(base_color, diff_rgb) + spec_rgb
# Apply Headlight
if wp.static(rc.headlight_active):
if wp.static(static_headlight_active):
cam_pos = ray_origin_world
cam_fwd = -cam_mat_world[:, 2]
hl_diff, hl_spec = compute_lighting(
@@ -891,15 +911,15 @@ def render(m: Model, d: Data, rc: RenderContext):
wp.vec3(1.0, 0.0, 0.0),
0.0,
0.0,
wp.static(rc.headlight_diffuse),
wp.static(rc.headlight_specular),
wp.static(static_headlight_diffuse),
wp.static(static_headlight_specular),
normal,
hit_point,
view_dir,
mat_spec,
mat_shin_exp,
wp.static(rc.enable_backface_culling),
wp.static(rc.enable_specular),
wp.static(static_enable_backface_culling),
wp.static(static_enable_specular),
True,
False,
)
+179
View File
@@ -0,0 +1,179 @@
# SPDX-FileCopyrightText: Copyright (c) 2023 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
# SPDX-License-Identifier: Apache-2.0
import warp
_wp_module_name_ = "warp.jax"
def device_to_jax(warp_device: warp.DeviceLike):
"""Return the Jax device corresponding to a Warp device.
Returns:
:class:`jax.Device`
Raises:
RuntimeError: Failed to find the corresponding Jax device.
"""
import jax # noqa: PLC0415
d = warp.get_device(warp_device)
if d.is_cuda:
cuda_devices = jax.devices("cuda")
if d.ordinal >= len(cuda_devices):
raise RuntimeError(f"Jax device corresponding to '{warp_device}' is not available")
return cuda_devices[d.ordinal]
else:
cpu_devices = jax.devices("cpu")
if not cpu_devices:
raise RuntimeError(f"Jax device corresponding to '{warp_device}' is not available")
return cpu_devices[0]
def device_from_jax(jax_device) -> warp._src.context.Device:
"""Return the Warp device corresponding to a Jax device.
Args:
jax_device (jax.Device): A Jax device descriptor.
Raises:
RuntimeError: The Jax device is neither a CPU nor GPU device.
"""
if jax_device.platform == "cpu":
return warp.get_device("cpu")
elif jax_device.platform == "gpu":
return warp.get_cuda_device(jax_device.id)
else:
raise RuntimeError(f"Unsupported Jax device platform '{jax_device.platform}'")
def get_jax_device():
"""Get the current Jax device."""
import jax # noqa: PLC0415
# TODO: is there a simpler way of getting the Jax "current" device?
# check if jax.default_device() context manager is active
device = jax.config.jax_default_device
# if default device is not set, use first device
if device is None:
device = jax.local_devices()[0]
return device
def dtype_to_jax(warp_dtype):
"""Return the Jax dtype corresponding to a Warp dtype.
Args:
warp_dtype: A Warp data type that has a corresponding Jax data type.
Raises:
TypeError: Unable to find a corresponding Jax data type.
"""
# initialize lookup table on first call to defer jax import
if dtype_to_jax.type_map is None:
import jax.numpy as jp # noqa: PLC0415
dtype_to_jax.type_map = {
warp.float16: jp.float16,
warp.bfloat16: jp.bfloat16,
warp.float32: jp.float32,
warp.float64: jp.float64,
warp.int8: jp.int8,
warp.int16: jp.int16,
warp.int32: jp.int32,
warp.int64: jp.int64,
warp.uint8: jp.uint8,
warp.uint16: jp.uint16,
warp.uint32: jp.uint32,
warp.uint64: jp.uint64,
warp.bool: jp.bool_,
}
jax_dtype = dtype_to_jax.type_map.get(warp_dtype)
if jax_dtype is not None:
return jax_dtype
else:
raise TypeError(f"Cannot convert {warp_dtype} to a Jax type")
def dtype_from_jax(jax_dtype):
"""Return the Warp dtype corresponding to a Jax dtype.
Raises:
TypeError: Unable to find a corresponding Warp data type.
"""
# initialize lookup table on first call to defer jax import
if dtype_from_jax.type_map is None:
import jax.numpy as jp # noqa: PLC0415
dtype_from_jax.type_map = {
# Jax scalar types
jp.float16: warp.float16,
jp.bfloat16: warp.bfloat16,
jp.float32: warp.float32,
jp.float64: warp.float64,
jp.int8: warp.int8,
jp.int16: warp.int16,
jp.int32: warp.int32,
jp.int64: warp.int64,
jp.uint8: warp.uint8,
jp.uint16: warp.uint16,
jp.uint32: warp.uint32,
jp.uint64: warp.uint64,
jp.bool_: warp.bool,
# Jax dtype objects
jp.dtype(jp.float16): warp.float16,
jp.dtype(jp.bfloat16): warp.bfloat16,
jp.dtype(jp.float32): warp.float32,
jp.dtype(jp.float64): warp.float64,
jp.dtype(jp.int8): warp.int8,
jp.dtype(jp.int16): warp.int16,
jp.dtype(jp.int32): warp.int32,
jp.dtype(jp.int64): warp.int64,
jp.dtype(jp.uint8): warp.uint8,
jp.dtype(jp.uint16): warp.uint16,
jp.dtype(jp.uint32): warp.uint32,
jp.dtype(jp.uint64): warp.uint64,
jp.dtype(jp.bool_): warp.bool,
}
wp_dtype = dtype_from_jax.type_map.get(jax_dtype)
if wp_dtype is not None:
return wp_dtype
else:
raise TypeError(f"Cannot convert {jax_dtype} to a Warp type")
# lookup tables initialized when needed
dtype_from_jax.type_map = None
dtype_to_jax.type_map = None
def to_jax(warp_array):
"""
Convert a Warp array to a Jax array without copying the data.
Args:
warp_array (warp.array): The Warp array to convert.
Returns:
jax.Array: The converted Jax array.
"""
import jax.dlpack # noqa: PLC0415
return jax.dlpack.from_dlpack(warp_array)
def from_jax(jax_array, dtype=None) -> warp.array:
"""Convert a Jax array to a Warp array without copying the data.
Args:
jax_array (jax.Array): The Jax array to convert.
dtype: The target data type of the resulting Warp array. Defaults to the Jax array's data type mapped to a Warp data type.
Returns:
warp.array: The converted Warp array.
"""
return warp.from_dlpack(jax_array, dtype=dtype)
@@ -5,12 +5,12 @@ import ctypes
from functools import reduce
import warp as wp
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, matches_array_class, strides_from_shape
from warp._src.utils import warn
from warp._src.context import _build_launch_bounds, type_str
from mujoco.mjx.third_party.warp._src.jax import get_jax_device
from warp._src.logger import log_warning
from warp._src.types import array_t, matches_array_class, strides_from_shape
_wp_module_name_ = "warp.jax_experimental.custom_call"
_wp_module_name_ = "warp.jax.custom_call"
_jax_warp_p = None
@@ -25,11 +25,11 @@ def jax_kernel(kernel, launch_dims=None, quiet=False):
.. deprecated:: 1.10.0
This version of ``jax_kernel()`` is deprecated for JAX >= 0.5.0 and is not supported
with JAX >= 0.8.0. Use :func:`warp.jax_experimental.ffi.jax_kernel` instead, which
with JAX >= 0.8.0. Use :func:`warp.jax_kernel` instead, which
is the default implementation as of Warp 1.10.
This implementation requires JAX version 0.4.25 - 0.7.x. For JAX 0.8.0 and later,
use the FFI-based implementation at :func:`warp.jax_experimental.ffi.jax_kernel`.
use the FFI-based implementation at :func:`warp.jax_kernel`.
Args:
kernel: The Warp kernel to be wrapped.
@@ -56,14 +56,14 @@ def jax_kernel(kernel, launch_dims=None, quiet=False):
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."
msg += " Please use warp.jax_kernel instead."
raise RuntimeError(msg)
# deprecation warning
if jax.__version_info__ >= (0, 5, 0) and not quiet:
warn(
log_warning(
"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). "
"Please use the newer FFI version instead (warp.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,
@@ -97,7 +97,7 @@ def _warp_custom_callback(stream, buffers, opaque, opaque_len):
# Parse launch dimensions.
dims = [int(d) for d in dim_str.split(",")]
bounds = launch_bounds_t(dims)
bounds = _build_launch_bounds(dims, kernel.adj.kernel_dim)
# Parse arguments.
arg_strings = args_str.split(";")
@@ -11,25 +11,22 @@ import traceback
from collections.abc import Callable
from enum import IntEnum
import jax
import warp as wp
from warp._src.codegen import get_full_arg_spec, make_full_qualified_name
from warp._src.context import CudaMemcpyKind
from warp._src.jax import get_jax_device
from warp._src.context import CudaMemcpyKind, _build_launch_bounds
from mujoco.mjx.third_party.warp._src.jax import get_jax_device
from warp._src.logger import log_warning
from warp._src.types import (
array_t,
launch_bounds_t,
matches_array_class,
strides_from_shape,
type_size_in_bytes,
type_to_warp,
)
from warp._src.utils import warn
from .xla_ffi import *
_wp_module_name_ = "warp.jax_experimental.ffi"
_wp_module_name_ = "warp.jax.ffi"
# Holders for the custom callbacks to keep them alive.
_FFI_KERNEL_REGISTRY: dict[tuple, FfiKernel] = {}
@@ -41,8 +38,20 @@ _FFI_REGISTRY_LOCK = threading.Lock()
# Lock when XLA invokes callbacks from multiple threads.
_FFI_CALLBACK_LOCK = threading.Lock()
# Sentinel for detecting when per-call kwargs are passed to differentiable wrappers.
_MISSING = object()
JAX_CALLABLE_DEFAULT_GRAPH_CACHE_MAX = 32
def _get_jax():
import jax # noqa: PLC0415
return jax
def check_jax_version():
jax = _get_jax()
# check if JAX version supports this
if jax.__version_info__ < (0, 5, 0):
msg = (
@@ -69,8 +78,8 @@ def compute_batch_size(shape, batch_ndim):
return batch_size
class GraphMode(IntEnum):
"""CUDA graph capture modes for :func:`warp.jax_experimental.jax_callable`.
class JaxCallableGraphMode(IntEnum):
"""CUDA graph capture modes for :func:`warp.jax_callable`.
These modes control whether JAX or Warp captures a CUDA graph, and whether
staging buffers are used when capturing with Warp.
@@ -88,12 +97,20 @@ class GraphMode(IntEnum):
"""Capture a Warp graph using staging buffers and perform memcpy outside the graph."""
class ModulePreloadMode(IntEnum):
GraphMode = JaxCallableGraphMode
class JaxModulePreloadMode(IntEnum):
"""Module preload modes for JAX interop callables."""
NONE = 0 # don't preload modules
CURRENT_DEVICE = 1 # preload on currently active device
ALL_DEVICES = 2 # preload on all supported devices
ModulePreloadMode = JaxModulePreloadMode
class FfiArg:
def __init__(self, name, type, in_out=False):
self.name = name
@@ -211,6 +228,7 @@ class FfiKernel:
self.input_output_aliases = input_output_aliases
# register the callback
jax = _get_jax()
FFI_CCALLFUNC = ctypes.CFUNCTYPE(
ctypes.c_void_p, ctypes.POINTER(XLA_FFI_CallFrame)
)
@@ -226,6 +244,7 @@ class FfiKernel:
jax.ffi.register_ffi_target(self.name, ffi_capsule_host, platform="Host")
def __call__(self, *args, output_dims=None, launch_dims=None, vmap_method=None):
jax = _get_jax()
num_inputs = len(args)
if num_inputs != self.num_inputs:
raise ValueError(f"Expected {self.num_inputs} inputs, but got {num_inputs}")
@@ -314,19 +333,19 @@ class FfiKernel:
)
# preload on the specified devices
if self.module_preload_mode == ModulePreloadMode.CURRENT_DEVICE:
if self.module_preload_mode == JaxModulePreloadMode.CURRENT_DEVICE:
device = wp.device_from_jax(get_jax_device())
self.kernel.module.load(device)
elif self.module_preload_mode == ModulePreloadMode.ALL_DEVICES:
elif self.module_preload_mode == JaxModulePreloadMode.ALL_DEVICES:
for d in jax.local_devices():
try:
dev = wp.device_from_jax(d)
# we only support CUDA devices for now
if dev.is_cuda:
self.kernel.module.load(dev)
except Exception:
# ignore unsupported devices like TPUs
pass
# 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
@@ -429,7 +448,7 @@ class FfiKernel:
# roll batch size into the first launch dimension
launch_dims = (batch_size * launch_dims[0], *launch_dims[1:])
launch_bounds = launch_bounds_t(launch_dims)
launch_bounds = _build_launch_bounds(launch_dims, self.kernel.adj.kernel_dim)
kernel_params[0] = ctypes.addressof(launch_bounds)
# get device and stream
@@ -500,7 +519,7 @@ class FfiCallDesc:
class FfiCallable:
default_graph_cache_max: int | None = 32
default_graph_cache_max: int | None = JAX_CALLABLE_DEFAULT_GRAPH_CACHE_MAX
def __init__(
self,
@@ -616,6 +635,7 @@ class FfiCallable:
self.input_output_aliases = input_output_aliases
# register the callback
jax = _get_jax()
FFI_CCALLFUNC = ctypes.CFUNCTYPE(
ctypes.c_void_p, ctypes.POINTER(XLA_FFI_CallFrame)
)
@@ -631,6 +651,7 @@ class FfiCallable:
jax.ffi.register_ffi_target(self.name, ffi_capsule_host, platform="Host")
def __call__(self, *args, output_dims=None, vmap_method=None):
jax = _get_jax()
num_inputs = len(args)
if num_inputs != self.num_inputs:
input_names = ", ".join(arg.name for arg in self.input_args)
@@ -709,18 +730,18 @@ class FfiCallable:
# preload on the specified devices
# NOTE: if the target function uses kernels from different modules, they will not be loaded here
module = wp.get_module(self.func.__module__)
if self.module_preload_mode == ModulePreloadMode.CURRENT_DEVICE:
if self.module_preload_mode == JaxModulePreloadMode.CURRENT_DEVICE:
device = wp.device_from_jax(get_jax_device())
module.load(device)
elif self.module_preload_mode == ModulePreloadMode.ALL_DEVICES:
elif self.module_preload_mode == JaxModulePreloadMode.ALL_DEVICES:
for d in jax.local_devices():
try:
dev = wp.device_from_jax(d)
if dev.is_cuda or dev.is_cpu:
module.load(dev)
except Exception:
# ignore unsupported devices like TPUs
pass
if dev.is_cuda or dev.is_cpu:
module.load(dev)
# save call data to be retrieved by callback
call_id = self.call_id
@@ -741,7 +762,7 @@ class FfiCallable:
metadata_ext.contents.metadata.contents.api_version.major_version = 0
metadata_ext.contents.metadata.contents.api_version.minor_version = 1
# Turn on CUDA graphs for this handler if on CUDA platform.
if self.graph_mode is GraphMode.JAX and platform == "CUDA":
if self.graph_mode is JaxCallableGraphMode.JAX and platform == "CUDA":
metadata_ext.contents.metadata.contents.traits = (
XLA_FFI_Handler_TraitsBits.COMMAND_BUFFER_COMPATIBLE
)
@@ -800,7 +821,7 @@ class FfiCallable:
cuda_stream = get_stream_from_callframe(call_frame.contents)
device_ordinal = get_device_ordinal_from_callframe(call_frame.contents)
if self.graph_mode == GraphMode.WARP:
if self.graph_mode == JaxCallableGraphMode.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]
@@ -819,7 +840,7 @@ class FfiCallable:
# early out
return
elif self.graph_mode == GraphMode.WARP_STAGED_EX:
elif self.graph_mode == JaxCallableGraphMode.WARP_STAGED_EX:
if call_desc.capture is not None:
graph_exec = call_desc.capture.graph.graph_exec
context = call_desc.capture.graph.device.context
@@ -866,7 +887,7 @@ class FfiCallable:
# early out
return
elif self.graph_mode == GraphMode.WARP_STAGED:
elif self.graph_mode == JaxCallableGraphMode.WARP_STAGED:
if call_desc.capture is not None:
graph_exec = call_desc.capture.graph.graph_exec
@@ -938,7 +959,7 @@ class FfiCallable:
# keep a reference to the capture object to prevent required modules getting unloaded
call_desc.capture = capture
elif self.graph_mode == GraphMode.WARP and device.is_cuda:
elif self.graph_mode == JaxCallableGraphMode.WARP and device.is_cuda:
# capturing with WARP
with wp.ScopedCapture() as capture:
self.func(*arg_list)
@@ -951,7 +972,7 @@ class FfiCallable:
if self._graph_cache_max is not None and len(self.captures) > self._graph_cache_max:
self.captures.popitem(last=False)
elif self.graph_mode == GraphMode.WARP_STAGED_EX and device.is_cuda:
elif self.graph_mode == JaxCallableGraphMode.WARP_STAGED_EX and device.is_cuda:
# capturing with WARP using staging buffers and memcopies done outside of the graph
wp_memcpy_batch = wp._src.context.runtime.core.wp_memcpy_batch
@@ -994,7 +1015,7 @@ class FfiCallable:
# TODO: we should have a way of freeing this
call_desc.capture = capture
elif self.graph_mode == GraphMode.WARP_STAGED and device.is_cuda:
elif self.graph_mode == JaxCallableGraphMode.WARP_STAGED and device.is_cuda:
# capturing with WARP using staging buffers and memcopies done inside of the graph
wp_cuda_graph_insert_memcpy_batch = (
wp._src.context.runtime.core.wp_cuda_graph_insert_memcpy_batch
@@ -1198,7 +1219,7 @@ def jax_kernel(
launch_dims=None,
output_dims=None,
in_out_argnames=None,
module_preload_mode=ModulePreloadMode.CURRENT_DEVICE,
module_preload_mode=JaxModulePreloadMode.CURRENT_DEVICE,
enable_backward: bool = False,
has_side_effect: bool = False,
):
@@ -1212,12 +1233,16 @@ def jax_kernel(
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.
launch_dims: Specify the kernel launch dimensions. If None, launch dimensions
are inferred from the shape of the first array argument. When
``enable_backward=False``, this value will be used by default but
can be overridden for individual calls. When ``enable_backward=True``,
this value is fixed at construction time and cannot be overridden
per call.
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.
Not supported when ``enable_backward=True``.
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``.
@@ -1233,9 +1258,11 @@ def jax_kernel(
- 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.
- ``output_dims`` and ``in_out_argnames`` are not supported when ``enable_backward=True``.
"""
check_jax_version()
jax = _get_jax()
if isinstance(output_dims, dict):
hashable_output_dims = tuple(sorted(output_dims.items()))
@@ -1283,11 +1310,11 @@ def jax_kernel(
"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."
)
# TODO: support output_dims with enable_backward=True (requires separate
# output-buffer allocation logic). launch_dims is supported below: the
# captured value is applied to both the forward and the adjoint launches.
if output_dims is not None:
raise NotImplementedError("jax_kernel(): output_dims is not yet supported when enable_backward=True.")
# Differentiable path: build a custom VJP wrapper inline.
# Infer the original kernel signature (names and annotations)
@@ -1307,8 +1334,17 @@ def jax_kernel(
else:
raise TypeError(f"Invalid type for argument '{p.name}', expected array or scalar, got {type}")
# Capture an explicit user-supplied launch_dims so the same value is used
# for both the forward launch and the adjoint launch (required for
# correct gradient values when array.ndim > kernel.tid_ndim).
# Reuse `hashable_launch_dims` (computed above for the cache key path)
# so 1-D integer and sequence forms are normalized identically.
_user_launch_dims = hashable_launch_dims if launch_dims is not None else None
def _resolve_launch_dims(call_args):
# determine launch dimensions from the shape of the first input array
if _user_launch_dims is not None:
return _user_launch_dims
# Fallback: determine launch dimensions from the shape of the first input array
for i, p in enumerate(parameters[:num_inputs]):
param_type = p.annotation
if matches_array_class(param_type, wp.array):
@@ -1360,11 +1396,15 @@ def jax_kernel(
try:
gi.zero_()
except Exception as e:
warn(f"Failed to zero gradient array: {e}", stacklevel=2)
log_warning(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.
# The same _resolve_launch_dims() is used here so that the adjoint
# kernel launches with exactly the same iteration space as the forward
# kernel (captured from the enclosing scope via _user_launch_dims when
# the caller supplied an explicit value, otherwise inferred from the
# inputs). This matches the forward path and avoids N x
# over-accumulation in atomic_add when array.ndim > kernel.tid_ndim.
wp.launch(
kernel,
dim=_resolve_launch_dims(inputs),
@@ -1498,46 +1538,65 @@ def jax_kernel(
vmap_method,
module_preload_mode,
has_side_effect,
# Include the normalized launch_dims so that wrapping the same kernel
# with different launch_dims produces independent cache entries.
# Reusing _user_launch_dims ensures int and 1-tuple forms of the same
# value map to the same key.
_user_launch_dims,
)
if static_args:
static_names = [parameters[i].name for i in static_args]
else:
static_names = []
def _user_callable(*args):
return jax_func(*args)
key = (*key, tuple(sorted(static_names)))
_user_callable.__signature__ = signature
def _user_callable(*args):
return jax_func(*args)
# Cache differentiable wrapper
key = (*key, 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]
_user_callable.__signature__ = signature
# Cache differentiable wrapper (no static args)
key = (*key, ())
# Cache differentiable wrapper
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
cached = jax.jit(_user_callable, static_argnames=tuple(static_names))
_FFI_DIFF_KERNEL_REGISTRY[key] = cached
# Thin Python-level wrapper that intercepts FfiKernel-style per-call kwargs
# and raises informative errors before JAX tracing begins.
def _checked_wrapper(*args, launch_dims=_MISSING, output_dims=_MISSING, vmap_method=_MISSING):
if launch_dims is not _MISSING:
raise TypeError(
"jax_kernel(): launch_dims cannot be overridden per-call when enable_backward=True "
f"(this wrapper was created with launch_dims={_user_launch_dims!r}). "
"Call jax_kernel() again with a different launch_dims to create a new wrapper."
)
if output_dims is not _MISSING:
raise TypeError("jax_kernel(): output_dims is not supported when enable_backward=True.")
if vmap_method is not _MISSING:
raise TypeError(
"jax_kernel(): vmap_method cannot be overridden per-call when enable_backward=True; "
"it is fixed at construction time. "
"Call jax_kernel() again with a different vmap_method to create a new wrapper."
)
return cached(*args)
return _checked_wrapper
def jax_callable(
func: Callable,
num_outputs: int = 1,
graph_mode: GraphMode = GraphMode.JAX,
graph_mode: JaxCallableGraphMode = JaxCallableGraphMode.JAX,
vmap_method: str | None = "broadcast_all",
output_dims=None,
in_out_argnames=None,
stage_in_argnames=None,
stage_out_argnames=None,
graph_cache_max: int | None = None,
module_preload_mode: ModulePreloadMode = ModulePreloadMode.CURRENT_DEVICE,
graph_cache_max: int | None = JAX_CALLABLE_DEFAULT_GRAPH_CACHE_MAX,
module_preload_mode: JaxModulePreloadMode = JaxModulePreloadMode.CURRENT_DEVICE,
has_side_effect: bool = False,
):
"""Create a JAX callback from an annotated Python function.
@@ -1551,10 +1610,10 @@ def jax_callable(
num_outputs: Specify the number of output arguments if greater than 1.
This must include the number of ``in_out_arguments``.
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,
``JaxCallableGraphMode.JAX`` (default): Let JAX capture the graph, which may be used as a subgraph in an enclosing JAX capture.
``JaxCallableGraphMode.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,
``JaxCallableGraphMode.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.
@@ -1564,12 +1623,12 @@ def jax_callable(
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``.
stage_in_argnames: Names of input arguments that need to be copied with ``GraphMode.WARP_STAGED*``.
stage_in_argnames: Names of input arguments that need to be copied with ``JaxCallableGraphMode.WARP_STAGED*``.
If ``None``, copy all input arguments.
stage_out_argnames: Names of output arguments that need to be copied with ``GraphMode.WARP_STAGED*``.
stage_out_argnames: Names of output arguments that need to be copied with ``JaxCallableGraphMode.WARP_STAGED*``.
If ``None``, copy all output arguments.
graph_cache_max: Maximum number of cached graphs captured using ``GraphMode.WARP``.
If ``None``, use ``warp.jax_experimental.get_jax_callable_default_graph_cache_max()``.
graph_cache_max: Maximum number of cached graphs captured using ``JaxCallableGraphMode.WARP``.
If ``None``, the graph cache is unlimited.
module_preload_mode: Specify the devices where the module should be preloaded.
has_side_effect: Whether the custom call has side effects. When True,
the FFI call will be executed even when the outputs are not used.
@@ -1584,9 +1643,6 @@ def jax_callable(
check_jax_version()
if graph_cache_max is None:
graph_cache_max = FfiCallable.default_graph_cache_max
if isinstance(output_dims, dict):
hashable_output_dims = tuple(sorted(output_dims.items()))
elif hasattr(output_dims, "__len__"):
@@ -1631,14 +1687,14 @@ def jax_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``.
Get the maximum size of the graph cache for graphs captured using ``JaxCallableGraphMode.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``.
Set the maximum size of the graph cache for graphs captured using ``JaxCallableGraphMode.WARP``, unlimited if ``None``.
"""
FfiCallable.default_graph_cache_max = cache_max
@@ -1677,6 +1733,7 @@ def register_ffi_callback(name: str, func: Callable, graph_compatible: bool = Tr
"""
check_jax_version()
jax = _get_jax()
# TODO check that the name is not already registered
@@ -1765,6 +1822,8 @@ def get_warp_shape(arg, dims):
def get_jax_output_type(arg, dims):
jax = _get_jax()
if isinstance(dims, int):
dims = (dims,)
@@ -3,13 +3,78 @@
import ctypes
import enum
import threading
import jax.numpy as jnp
import numpy as np
import warp as wp
_wp_module_name_ = "warp.jax_experimental.xla_ffi"
_wp_module_name_ = "warp.jax.xla_ffi"
_xla_data_type_to_constructor = None
_XLA_DATA_TYPE_TO_CONSTRUCTOR_LOCK = threading.Lock()
def _get_jnp():
import jax.numpy as jnp # noqa: PLC0415
return jnp
def _get_xla_data_type_to_constructor():
global _xla_data_type_to_constructor
if _xla_data_type_to_constructor is None:
with _XLA_DATA_TYPE_TO_CONSTRUCTOR_LOCK:
if _xla_data_type_to_constructor is None:
jnp = _get_jnp()
constructors = {
# 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"):
constructors[XLA_FFI_DataType.F8E5M2] = jnp.float8_e5m2
if hasattr(jnp, "float8_e3m4"):
constructors[XLA_FFI_DataType.F8E3M4] = jnp.float8_e3m4
if hasattr(jnp, "float8_e4m3"):
constructors[XLA_FFI_DataType.F8E4M3] = jnp.float8_e4m3
if hasattr(jnp, "float8_e4m3fn"):
constructors[XLA_FFI_DataType.F8E4M3FN] = jnp.float8_e4m3fn
if hasattr(jnp, "float8_e4m3b11fnuz"):
constructors[XLA_FFI_DataType.F8E4M3B11FNUZ] = jnp.float8_e4m3b11fnuz
if hasattr(jnp, "float8_e5m2fnuz"):
constructors[XLA_FFI_DataType.F8E5M2FNUZ] = jnp.float8_e5m2fnuz
if hasattr(jnp, "float8_e4m3fnuz"):
constructors[XLA_FFI_DataType.F8E4M3FNUZ] = jnp.float8_e4m3fnuz
_xla_data_type_to_constructor = constructors
return _xla_data_type_to_constructor
def _jnp_dtype_for_xla(xla_dtype):
jnp = _get_jnp()
return jnp.dtype(_get_xla_data_type_to_constructor()[xla_dtype])
#######################################################################
# ctypes structures and enums for XLA's FFI API:
@@ -470,45 +535,6 @@ class XLA_FFI_CallFrame(ctypes.Structure):
)
_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
#######################################################################
@@ -522,14 +548,14 @@ def decode_bytespan(span: XLA_FFI_ByteSpan):
def decode_scalar(scalar: XLA_FFI_Scalar):
# TODO validate if dtype supported
dtype = jnp.dtype(_xla_data_type_to_constructor[scalar.dtype])
dtype = _jnp_dtype_for_xla(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])
dtype = _jnp_dtype_for_xla(array.dtype)
bytes = ctypes.string_at(array.data, dtype.itemsize * array.size)
return np.frombuffer(bytes, dtype=dtype)
@@ -612,7 +638,7 @@ def dtype_from_ffi(ffi_dtype):
def jax_dtype_from_ffi(ffi_dtype):
return _xla_data_type_to_constructor.get(ffi_dtype)
return _get_xla_data_type_to_constructor().get(ffi_dtype)
# Execution context (stream, stage)
@@ -632,7 +658,7 @@ class FfiBuffer:
def __init__(self, xla_buffer):
# TODO check if valid
self.dtype = jnp.dtype(_xla_data_type_to_constructor[xla_buffer.dtype])
self.dtype = _jnp_dtype_for_xla(xla_buffer.dtype)
self.shape = tuple(xla_buffer.dims[i] for i in range(xla_buffer.rank))
self.data = xla_buffer.data
@@ -1,2 +0,0 @@
# SPDX-FileCopyrightText: Copyright (c) 2025 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
# SPDX-License-Identifier: Apache-2.0
+4 -7
View File
@@ -14,19 +14,17 @@
# ==============================================================================
"""DO NOT EDIT. This file is auto-generated."""
import dataclasses
import functools
import jax
import warp as wp
from mujoco.mjx._src import types
import mujoco.mjx.third_party.mujoco_warp as mjwarp
from mujoco.mjx.third_party.mujoco_warp._src import types as mjwp_types
from mujoco.mjx.warp import ffi
from mujoco.mjx.warp.render_context import _MJX_RENDER_CONTEXT_BUFFERS
from mujoco.mjx.warp.render_context import RenderContextPytree
import mujoco.mjx.third_party.mujoco_warp as mjwarp
from mujoco.mjx.third_party.mujoco_warp._src import types as mjwp_types
import warp as wp
_m = mjwarp.Model(
**{f.name: None for f in dataclasses.fields(mjwarp.Model) if f.init}
@@ -50,7 +48,6 @@ _cb = mjwp_types.Callback(
**{f.name: None for f in dataclasses.fields(mjwp_types.Callback) if f.init}
)
@ffi.format_args_for_warp
def _refit_bvh_shim(
# Model
+2 -6
View File
@@ -14,17 +14,14 @@
# ==============================================================================
"""DO NOT EDIT. This file is auto-generated."""
import dataclasses
import functools
import jax
import warp as wp
from mujoco.mjx._src import types
from mujoco.mjx.warp import ffi
import mujoco.mjx.third_party.mujoco_warp as mjwarp
from mujoco.mjx.third_party.mujoco_warp._src import types as mjwp_types
from mujoco.mjx.warp import ffi
import warp as wp
_m = mjwarp.Model(
**{f.name: None for f in dataclasses.fields(mjwarp.Model) if f.init}
@@ -48,7 +45,6 @@ _cb = mjwp_types.Callback(
**{f.name: None for f in dataclasses.fields(mjwp_types.Callback) if f.init}
)
@ffi.format_args_for_warp
def _collision_shim(
# Model
+5 -4
View File
@@ -22,10 +22,11 @@ from typing import Any, Callable, Optional, Sequence, Tuple, Union
import jax
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._src.jax_experimental import ffi
from mujoco.mjx.third_party.warp._src.jax import ffi as warp_ffi
from mujoco.mjx.warp import types as mjx_warp_types
def flatten_signature(signature: inspect.Signature, args: Tuple[Any, ...]):
@@ -97,7 +98,7 @@ def flatten_signature(signature: inspect.Signature, args: Tuple[Any, ...]):
def jax_callable_variadic_tuple(
func: Callable, # pylint: disable=g-bare-generic
num_outputs: int = 1,
graph_mode: ffi.GraphMode = ffi.GraphMode.WARP,
graph_mode: warp_ffi.JaxCallableGraphMode = warp_ffi.JaxCallableGraphMode.WARP,
vmap_method: Optional[str] = None,
output_dims: Optional[dict[str, tuple[int, ...]]] = None,
in_out_argnames: Optional[Sequence[str]] = None,
@@ -126,7 +127,7 @@ def jax_callable_variadic_tuple(
if new_signature.return_annotation is not inspect.Signature.empty:
func_wrapper.__annotations__['return'] = new_signature.return_annotation
my_callable = ffi.jax_callable(
my_callable = warp_ffi.jax_callable(
func_wrapper,
num_outputs=num_outputs,
graph_mode=graph_mode,
+5 -5
View File
@@ -23,6 +23,7 @@ import mujoco.mjx.third_party.mujoco_warp as mjwarp
from mujoco.mjx.third_party.mujoco_warp._src import types as mjwp_types
import warp as wp
_m = mjwarp.Model(
**{f.name: None for f in dataclasses.fields(mjwarp.Model) if f.init}
)
@@ -45,7 +46,6 @@ _cb = mjwp_types.Callback(
**{f.name: None for f in dataclasses.fields(mjwp_types.Callback) if f.init}
)
@ffi.format_args_for_warp
def _forward_shim(
# Model
@@ -1609,8 +1609,8 @@ def _forward_jax_impl(m: types.Model, d: types.Data):
m._impl.M_mulm_col,
m._impl.M_mulm_madr,
m._impl.M_mulm_rowadr,
m._impl.M_rowadr,
m._impl.M_rownnz,
m.M_rowadr,
m.M_rownnz,
m._impl.M_tiles,
m.actuator_acc0,
m.actuator_actadr,
@@ -3899,8 +3899,8 @@ def _step_jax_impl(m: types.Model, d: types.Data):
m._impl.M_mulm_col,
m._impl.M_mulm_madr,
m._impl.M_mulm_rowadr,
m._impl.M_rowadr,
m._impl.M_rownnz,
m.M_rowadr,
m.M_rownnz,
m._impl.M_tiles,
m.actuator_acc0,
m.actuator_actadr,
+4 -7
View File
@@ -14,19 +14,17 @@
# ==============================================================================
"""DO NOT EDIT. This file is auto-generated."""
import dataclasses
import functools
import jax
import warp as wp
from mujoco.mjx._src import types
import mujoco.mjx.third_party.mujoco_warp as mjwarp
from mujoco.mjx.third_party.mujoco_warp._src import types as mjwp_types
from mujoco.mjx.warp import ffi
from mujoco.mjx.warp.render_context import _MJX_RENDER_CONTEXT_BUFFERS
from mujoco.mjx.warp.render_context import RenderContextPytree
import mujoco.mjx.third_party.mujoco_warp as mjwarp
from mujoco.mjx.third_party.mujoco_warp._src import types as mjwp_types
import warp as wp
_m = mjwarp.Model(
**{f.name: None for f in dataclasses.fields(mjwarp.Model) if f.init}
@@ -50,7 +48,6 @@ _cb = mjwp_types.Callback(
**{f.name: None for f in dataclasses.fields(mjwp_types.Callback) if f.init}
)
@ffi.format_args_for_warp
def _render_shim(
# Model
+25 -7
View File
@@ -14,6 +14,7 @@
# ==============================================================================
import functools
import os
import tempfile
from absl.testing import absltest
from absl.testing import parameterized
@@ -70,11 +71,21 @@ def _get_model_data_rc(xml, batch_size, render_seg=False):
class RenderTest(parameterized.TestCase):
@classmethod
def setUpClass(cls):
super().setUpClass()
if mjxw.WARP_INSTALLED:
cls.tempdir = tempfile.TemporaryDirectory()
wp.config.kernel_cache_dir = cls.tempdir.name
@classmethod
def tearDownClass(cls):
super().tearDownClass()
if hasattr(cls, 'tempdir'):
cls.tempdir.cleanup()
def setUp(self):
super().setUp()
if mjxw.WARP_INSTALLED:
tempdir = '/tmp/wp_kernel_cache_dir_RenderTest'
wp.config.kernel_cache_dir = tempdir
np.random.seed(0)
def _maybe_skip(self):
@@ -251,11 +262,18 @@ class RenderTest(parameterized.TestCase):
class RenderContextGarbageCollectionTest(absltest.TestCase):
"""Tests that RenderContext cleans up buffers on deletion."""
def setUp(self):
super().setUp()
@classmethod
def setUpClass(cls):
super().setUpClass()
if mjxw.WARP_INSTALLED:
tempdir = '/tmp/wp_kernel_cache_dir_RenderContextGCTest'
wp.config.kernel_cache_dir = tempdir
cls.tempdir = tempfile.TemporaryDirectory()
wp.config.kernel_cache_dir = cls.tempdir.name
@classmethod
def tearDownClass(cls):
super().tearDownClass()
if hasattr(cls, 'tempdir'):
cls.tempdir.cleanup()
def _maybe_skip(self):
if not mjxw.WARP_INSTALLED:
+2 -6
View File
@@ -14,17 +14,14 @@
# ==============================================================================
"""DO NOT EDIT. This file is auto-generated."""
import dataclasses
import functools
import jax
import warp as wp
from mujoco.mjx._src import types
from mujoco.mjx.warp import ffi
import mujoco.mjx.third_party.mujoco_warp as mjwarp
from mujoco.mjx.third_party.mujoco_warp._src import types as mjwp_types
from mujoco.mjx.warp import ffi
import warp as wp
_m = mjwarp.Model(
**{f.name: None for f in dataclasses.fields(mjwarp.Model) if f.init}
@@ -48,7 +45,6 @@ _cb = mjwp_types.Callback(
**{f.name: None for f in dataclasses.fields(mjwp_types.Callback) if f.init}
)
@ffi.format_args_for_warp
def _kinematics_shim(
# Model
+2 -2
View File
@@ -32,10 +32,10 @@ from mujoco.mjx.warp import collision_driver as wp_collision
from mujoco.mjx.warp import forward as wp_forward
from mujoco.mjx.warp import io
from mujoco.mjx.warp import smooth as wp_smooth
from mujoco.mjx.warp import types as mjxw_types
import mujoco.mjx.third_party.mujoco_warp as mjwarp
import numpy as np
import warp as wp
from mujoco.mjx.third_party.warp._src.jax_experimental import ffi as warp_ffi
_MODELFILE = flags.DEFINE_string(
@@ -280,7 +280,7 @@ def _main(_: Sequence[str]):
m = mujoco.MjModel.from_xml_path(modelfile)
benchmark_type = _BENCHMARK.value
graph_mode = getattr(warp_ffi.GraphMode, _GRAPH_MODE.value)
graph_mode = getattr(mjxw_types.GraphMode, _GRAPH_MODE.value)
# Only allocate the model needed for the specific benchmark
mx = None
+3 -5
View File
@@ -23,7 +23,6 @@ from jax import tree_util
from jax.interpreters import batching
from mujoco.mjx._src import dataclasses as mjx_dataclasses
import numpy as np
if typing.TYPE_CHECKING:
GraphMode = int
@@ -33,7 +32,9 @@ if typing.TYPE_CHECKING:
else:
try:
from warp._src.jax_experimental.ffi import GraphMode
from mujoco.mjx.third_party.warp._src.jax import ffi as warp_ffi
GraphMode = warp_ffi.JaxCallableGraphMode
from mujoco.mjx.third_party.mujoco_warp._src import types as mjwp_types
Callback = mjwp_types.Callback
@@ -170,7 +171,6 @@ class ModelWarp(PyTreeNode):
D_diag: np.ndarray
D_rowadr: np.ndarray
D_rownnz: np.ndarray
M_colind: np.ndarray
M_elemid: np.ndarray
M_fullm_i: np.ndarray
M_fullm_j: np.ndarray
@@ -180,8 +180,6 @@ class ModelWarp(PyTreeNode):
M_mulm_col: np.ndarray
M_mulm_madr: np.ndarray
M_mulm_rowadr: np.ndarray
M_rowadr: np.ndarray
M_rownnz: np.ndarray
M_tiles: Tuple[TileSet, ...]
actuator_delay: np.ndarray
actuator_history: np.ndarray
+1 -1
View File
@@ -37,7 +37,7 @@ dependencies = [
[project.optional-dependencies]
warp = [
"warp-lang==1.13.0",
"warp-lang==1.14.0",
]
dev = [
"isort",