Merge pull request #3465 from shi-eric:shi-eric/mjx-warp-cache-ffi-callables
PiperOrigin-RevId: 961136122 Change-Id: Ia468812e0557bb4efb7866f7130dab5d1cea673f
This commit is contained in:
+93
-27
@@ -17,6 +17,7 @@
|
||||
import dataclasses
|
||||
import functools
|
||||
import inspect
|
||||
import threading
|
||||
import typing
|
||||
from typing import Any, Callable, Optional, Sequence, Tuple, Union
|
||||
|
||||
@@ -30,6 +31,15 @@ from mujoco.mjx._src.types import tree_path_to_attr_str
|
||||
from mujoco.mjx.warp import types as mjx_warp_types
|
||||
|
||||
|
||||
# ``warp_ffi.jax_callable`` keys its registry by wrapper function and
|
||||
# configuration. Cache generated wrappers here by the original shim and MJX's
|
||||
# flattened call structure so equivalent retraces reuse the same Warp entry.
|
||||
_JAX_CALLABLE_VARIADIC_TUPLE_REGISTRY: dict[
|
||||
tuple[Any, ...], Callable[..., Any]
|
||||
] = {}
|
||||
_JAX_CALLABLE_VARIADIC_TUPLE_REGISTRY_LOCK = threading.Lock()
|
||||
|
||||
|
||||
def flatten_signature(signature: inspect.Signature, args: Tuple[Any, ...]):
|
||||
"""Flattens a tuple/dataclass signature."""
|
||||
|
||||
@@ -109,38 +119,94 @@ def jax_callable_variadic_tuple(
|
||||
):
|
||||
"""Wraps a JAX callable to support variadic tuples and dataclasses."""
|
||||
|
||||
# Snapshot mapping and name-set options into stable cache-key forms. Warp
|
||||
# consumes them by argument name, so caller ordering does not distinguish
|
||||
# callable configurations.
|
||||
hashable_output_dims = (
|
||||
None if output_dims is None else tuple(sorted(output_dims.items()))
|
||||
)
|
||||
hashable_in_out_argnames = (
|
||||
tuple(sorted(in_out_argnames)) if in_out_argnames else None
|
||||
)
|
||||
hashable_stage_in_argnames = (
|
||||
tuple(sorted(stage_in_argnames)) if stage_in_argnames else None
|
||||
)
|
||||
hashable_stage_out_argnames = (
|
||||
tuple(sorted(stage_out_argnames)) if stage_out_argnames else None
|
||||
)
|
||||
|
||||
def callable_wrapper(*args, **kwargs):
|
||||
def func_wrapper(*flat_args, **kwargs):
|
||||
num_inputs = in_tree.num_leaves
|
||||
flat_inputs = flat_args[:num_inputs]
|
||||
output_buffers = flat_args[num_inputs:]
|
||||
unflat_args = jax.tree.unflatten(in_tree, flat_inputs)
|
||||
return func(*unflat_args, *output_buffers, **kwargs)
|
||||
|
||||
# Provide a flattened signature for the Warp callable machinery.
|
||||
flat_args, in_tree = jax.tree.flatten(args)
|
||||
new_signature = flatten_signature(inspect.signature(func), args)
|
||||
func_wrapper.__signature__ = new_signature # pyrefly: ignore[missing-attribute]
|
||||
func_wrapper.__annotations__ = {
|
||||
p.name: p.annotation
|
||||
for p in new_signature.parameters.values()
|
||||
if p.annotation is not inspect.Parameter.empty
|
||||
}
|
||||
if new_signature.return_annotation is not inspect.Signature.empty:
|
||||
func_wrapper.__annotations__['return'] = new_signature.return_annotation
|
||||
|
||||
my_callable = warp_ffi.jax_callable(
|
||||
func_wrapper,
|
||||
num_outputs=num_outputs,
|
||||
graph_mode=graph_mode,
|
||||
vmap_method=vmap_method,
|
||||
output_dims=output_dims,
|
||||
in_out_argnames=in_out_argnames,
|
||||
stage_in_argnames=stage_in_argnames,
|
||||
stage_out_argnames=stage_out_argnames,
|
||||
has_side_effect=has_side_effect,
|
||||
# Cache the wrapper's structural ABI, not per-call leaves. The flattened
|
||||
# signature defines Warp's arguments and the PyTree defines reconstruction.
|
||||
# Leaf arrays and tracers are forwarded on every invocation; keying on them
|
||||
# would prevent equivalent retraces from reusing the same FFI target.
|
||||
callable_cache_key = (
|
||||
func,
|
||||
new_signature,
|
||||
in_tree,
|
||||
num_outputs,
|
||||
graph_mode,
|
||||
vmap_method,
|
||||
hashable_output_dims,
|
||||
hashable_in_out_argnames,
|
||||
hashable_stage_in_argnames,
|
||||
hashable_stage_out_argnames,
|
||||
has_side_effect,
|
||||
)
|
||||
|
||||
flat_args, in_tree = jax.tree.flatten(args)
|
||||
# Serialize construction so concurrent traces cannot register duplicate
|
||||
# Warp callbacks for the same structural key.
|
||||
with _JAX_CALLABLE_VARIADIC_TUPLE_REGISTRY_LOCK:
|
||||
my_callable = _JAX_CALLABLE_VARIADIC_TUPLE_REGISTRY.get(
|
||||
callable_cache_key
|
||||
)
|
||||
if my_callable is None:
|
||||
|
||||
# Restore the original PyTree inputs; Warp appends output buffers.
|
||||
def func_wrapper(*flat_args, **kwargs):
|
||||
num_inputs = in_tree.num_leaves
|
||||
flat_inputs = flat_args[:num_inputs]
|
||||
output_buffers = flat_args[num_inputs:]
|
||||
unflat_args = jax.tree.unflatten(in_tree, flat_inputs)
|
||||
return func(*unflat_args, *output_buffers, **kwargs)
|
||||
|
||||
# Warp derives the FFI ABI from this synthetic signature and
|
||||
# annotations.
|
||||
func_wrapper.__signature__ = ( # pyrefly: ignore[missing-attribute]
|
||||
new_signature
|
||||
)
|
||||
func_wrapper.__annotations__ = {
|
||||
p.name: p.annotation
|
||||
for p in new_signature.parameters.values()
|
||||
if p.annotation is not inspect.Parameter.empty
|
||||
}
|
||||
if new_signature.return_annotation is not inspect.Signature.empty:
|
||||
func_wrapper.__annotations__['return'] = (
|
||||
new_signature.return_annotation
|
||||
)
|
||||
|
||||
# Constructing the callable registers its FFI target.
|
||||
my_callable = warp_ffi.jax_callable(
|
||||
func_wrapper,
|
||||
num_outputs=num_outputs,
|
||||
graph_mode=graph_mode,
|
||||
vmap_method=vmap_method,
|
||||
output_dims=(
|
||||
None
|
||||
if hashable_output_dims is None
|
||||
else dict(hashable_output_dims)
|
||||
),
|
||||
in_out_argnames=hashable_in_out_argnames,
|
||||
stage_in_argnames=hashable_stage_in_argnames,
|
||||
stage_out_argnames=hashable_stage_out_argnames,
|
||||
has_side_effect=has_side_effect,
|
||||
)
|
||||
_JAX_CALLABLE_VARIADIC_TUPLE_REGISTRY[callable_cache_key] = my_callable
|
||||
|
||||
# Invocation may re-enter JAX or Warp and needs no registry serialization.
|
||||
return my_callable(*flat_args, **kwargs)
|
||||
|
||||
return callable_wrapper
|
||||
|
||||
@@ -0,0 +1,76 @@
|
||||
# Copyright 2026 DeepMind Technologies Limited
|
||||
#
|
||||
# Licensed under the Apache License, Version 2.0 (the "License");
|
||||
# you may not use this file except in compliance with the License.
|
||||
# You may obtain a copy of the License at
|
||||
#
|
||||
# http://www.apache.org/licenses/LICENSE-2.0
|
||||
#
|
||||
# Unless required by applicable law or agreed to in writing, software
|
||||
# distributed under the License is distributed on an "AS IS" BASIS,
|
||||
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
# See the License for the specific language governing permissions and
|
||||
# limitations under the License.
|
||||
# ==============================================================================
|
||||
"""Tests for MJX-Warp FFI helpers."""
|
||||
|
||||
import os
|
||||
from unittest import mock
|
||||
|
||||
from absl.testing import absltest
|
||||
from mujoco.mjx._src import io
|
||||
import mujoco.mjx.warp as mjxw
|
||||
|
||||
|
||||
try:
|
||||
from mujoco.mjx.warp import ffi # pylint: disable=g-import-not-at-top
|
||||
except ImportError:
|
||||
ffi = None
|
||||
|
||||
_FORCE_TEST = os.environ.get('MJX_WARP_FORCE_TEST', '0') == '1'
|
||||
|
||||
|
||||
class FfiTest(absltest.TestCase):
|
||||
|
||||
def test_jax_callable_variadic_tuple_cache(self):
|
||||
if not _FORCE_TEST:
|
||||
if not mjxw.WARP_INSTALLED:
|
||||
self.skipTest('Warp not installed.')
|
||||
if not io.has_cuda_gpu_device():
|
||||
self.skipTest('No CUDA GPU device available.')
|
||||
|
||||
def func_a(values: tuple[int, ...], output: int):
|
||||
del values, output
|
||||
|
||||
def func_b(values: tuple[int, ...], output: int):
|
||||
del values, output
|
||||
|
||||
# Count callback registrations without creating real JAX FFI targets.
|
||||
def create_callable(*args, **kwargs):
|
||||
del args, kwargs
|
||||
return mock.Mock()
|
||||
|
||||
with mock.patch.object(ffi, '_JAX_CALLABLE_VARIADIC_TUPLE_REGISTRY', {}):
|
||||
with mock.patch.object(
|
||||
ffi.warp_ffi, 'jax_callable', side_effect=create_callable
|
||||
) as jax_callable:
|
||||
wrapper_a = ffi.jax_callable_variadic_tuple(func_a)
|
||||
wrapper_a_again = ffi.jax_callable_variadic_tuple(func_a)
|
||||
wrapper_b = ffi.jax_callable_variadic_tuple(func_b)
|
||||
|
||||
# Equivalent traces of the same callable reuse one callback.
|
||||
wrapper_a((1, 2))
|
||||
wrapper_a_again((3, 4))
|
||||
self.assertEqual(jax_callable.call_count, 1)
|
||||
|
||||
# Distinct shim functions with the same signature do not share a target.
|
||||
wrapper_b((5, 6))
|
||||
self.assertEqual(jax_callable.call_count, 2)
|
||||
|
||||
# A different tuple arity changes the signature and PyTree cache key.
|
||||
wrapper_b((1, 2, 3))
|
||||
self.assertEqual(jax_callable.call_count, 3)
|
||||
|
||||
|
||||
if __name__ == '__main__':
|
||||
absltest.main()
|
||||
Reference in New Issue
Block a user