Implementation:Google deepmind Mujoco MJX Warp FFI: Difference between revisions
Auto-imported from implementations/Google_deepmind_Mujoco_MJX_Warp_FFI.md |
Sync from local file |
||
| Line 76: | Line 76: | ||
== Related Pages == | == Related Pages == | ||
* [[Google_deepmind_Mujoco_MJX_Warp_Collision_Driver]] - Uses @ffi.format_args_for_warp for collision shim | * [[Implementation:Google_deepmind_Mujoco_MJX_Warp_Collision_Driver]] - Uses @ffi.format_args_for_warp for collision shim | ||
* [[Google_deepmind_Mujoco_MJX_Warp_Forward]] - Uses @ffi.format_args_for_warp for forward dynamics shim | * [[Implementation:Google_deepmind_Mujoco_MJX_Warp_Forward]] - Uses @ffi.format_args_for_warp for forward dynamics shim | ||
* [[Google_deepmind_Mujoco_MJX_Warp_Smooth]] - Uses @ffi.format_args_for_warp for smooth dynamics shim | * [[Implementation:Google_deepmind_Mujoco_MJX_Warp_Smooth]] - Uses @ffi.format_args_for_warp for smooth dynamics shim | ||
* [[Google_deepmind_Mujoco_JAX_Experimental_FFI]] - Lower-level JAX-Warp FFI implementation used by this module | * [[Implementation:Google_deepmind_Mujoco_JAX_Experimental_FFI]] - Lower-level JAX-Warp FFI implementation used by this module | ||
[[Category:Implementations]] | [[Category:Implementations]] | ||
Latest revision as of 10:39, 27 September 2026
| Knowledge Sources | |
|---|---|
| Domains | Foreign Function Interface, JAX-Warp Bridge, GPU Computing, Type Marshalling |
| Last Updated | 2026-02-15 04:00 GMT |
Overview
FFI helper module that marshals JAX arrays and Python dataclasses into flattened argument signatures compatible with NVIDIA Warp kernel launches.
Description
ffi.py provides the core interoperability layer between JAX and Warp within the MJX pipeline. Its primary function flatten_signature takes an inspect.Signature and a tuple of arguments, then expands tuples and dataclasses into flat parameter lists that Warp kernels can consume. The @format_args_for_warp decorator applies this flattening logic transparently so that MJX shim functions (collision_driver, forward, smooth) can declare rich typed signatures while the FFI layer handles the conversion to Warp-compatible flat arrays.
Usage
This module is used by all MJX Warp bridge modules (collision_driver, forward, smooth) via the @ffi.format_args_for_warp decorator. It is also imported directly when custom Warp kernel invocations are needed from JAX code.
Code Reference
Source Location
- Repository: Google_deepmind_Mujoco
- File: mjx/mujoco/mjx/warp/ffi.py
- Lines: 1-418
Key Functions
def flatten_signature(signature: inspect.Signature, args: Tuple[Any, ...]):
"""Flattens a tuple/dataclass signature."""
def expand_parameter(parameter, arg_iter):
"""Expands a single parameter into flat sub-parameters for tuples and dataclasses."""
def format_args_for_warp(func):
"""Decorator that transforms rich-typed arguments into flat Warp-compatible arrays."""
Import
from mujoco.mjx.warp import ffi
from mujoco.mjx.warp.ffi import flatten_signature, format_args_for_warp
I/O Contract
Inputs
| Name | Type | Required | Description |
|---|---|---|---|
| signature | inspect.Signature | Yes | The function signature to flatten |
| args | Tuple[Any, ...] | Yes | Positional arguments containing JAX arrays, tuples, or dataclasses |
| func | Callable | Yes | The shim function to decorate (for format_args_for_warp) |
Outputs
| Name | Type | Description |
|---|---|---|
| flattened_params | list[inspect.Parameter] | Expanded flat parameter list with unique names (e.g., name__0, name__field) |
| wrapped_func | Callable | Decorated function that auto-flattens arguments before Warp dispatch |
Related Pages
- Implementation:Google_deepmind_Mujoco_MJX_Warp_Collision_Driver - Uses @ffi.format_args_for_warp for collision shim
- Implementation:Google_deepmind_Mujoco_MJX_Warp_Forward - Uses @ffi.format_args_for_warp for forward dynamics shim
- Implementation:Google_deepmind_Mujoco_MJX_Warp_Smooth - Uses @ffi.format_args_for_warp for smooth dynamics shim
- Implementation:Google_deepmind_Mujoco_JAX_Experimental_FFI - Lower-level JAX-Warp FFI implementation used by this module