Warp 互操作性实战指南:在 NumPy、PyTorch、JAX、Paddle 与 DLPack 之间零拷贝共享 GPU 数据 📅 发布时间:2026/9/17 3:11:15 👁 浏览次数: Warp 互操作性实战指南在 NumPy、PyTorch、JAX、Paddle 与 DLPack 之间零拷贝共享 GPU 数据【免费下载链接】warpA Python framework for GPU-accelerated simulation, robotics, and machine learning.项目地址: https://gitcode.com/GitHub_Trending/warp/warpWarpwarpA Python framework for GPU-accelerated simulation, robotics, and machine learning通过标准数组接口协议__array_interface__、__cuda_array_interface__、DLPack与 NumPy、CuPy、PyTorch、JAX、Paddle 等 Python 框架无缝互通。本文以 docs/user_guide/interoperability.rst 为主线系统讲解各类框架的数组转换、流stream同步、自动微分autograd接入与 DLPack 底层协议并结合仓库源码剖析转换函数的零拷贝实现原理。读完本文你将掌握如何在wp.launch中直接传入外部框架数组、如何用wp.from_*/wp.to_*系列函数做零拷贝互转、如何让 Warp 内核在 PyTorch / JAX 的自动微分与图捕获CUDA Graph体系中工作以及何时该绕过转换函数直接使用协议层。互操作性总览标准数组协议与快速参考Warp 与其他 Python 框架的互操作建立在标准接口协议之上。只要外部数组实现了以下任意一种协议就可以直接作为wp.launch的输入__array_interface__CPU 数组协议NumPy 等框架实现__cuda_array_interface__GPU 数组协议CuPy、PyTorch、Numba 等框架实现__dlpack__/__dlpack_device__DLPack 协议Python Array API 标准 v2022.12JAX、PyTorch、Paddle 均支持。从源码结构看协议支持被集中封装在 warp/_src/context.pyfrom_numpy位于约第 10013 行、warp/_src/torch.py、warp/_src/paddle.py、warp/_src/jax/init.py 与 warp/_src/dlpack.py 中形成统一的转换入口。各框架快速参考框架转换方式零拷贝梯度感知NumPywp.from_numpy/array.numpy()仅 CPU否PyTorchwp.from_torch/wp.to_torch是是JAXwp.from_jax/wp.to_jax是通过jax_kernelPaddlewp.from_paddle/wp.to_paddle是是CuPy / Numba__cuda_array_interface__协议是否DLPackwp.from_dlpack/framework.from_dlpack是否上表中梯度感知为Yes的框架其转换函数会在可用时把 Warp 的梯度数组与对应框架的 autograd 张量互转从而让 Warp 数组参与对方框架的反向传播计算。JAX 的梯度感知能力通过jax_kernelFFI包装器提供而非数组转换本身。直接传递数组最快的上手路径任何实现了__array_interface__CPU或__cuda_array_interface__GPU的对象都可以不调用任何转换函数直接传入wp.launch的inputs。这是大多数场景下最快的接入方式——省去了创建 Warp 数组对象的 CPU 开销。CPU 端示例NumPy 数组直接驱动 saxpy 内核import numpy as np import warp as wp wp.kernel def saxpy(x: wp.array[float], y: wp.array[float], a: float): i wp.tid() y[i] a * x[i] y[i] x np.arange(n, dtypenp.float32) y np.ones(n, dtypenp.float32) wp.launch(saxpy, dimn, inputs[x, y, 1.0], devicecpu)CUDA 端同样的模式适用于 CuPy、PyTorch 或任何暴露__cuda_array_interface__的框架import cupy as cp with cp.cuda.Device(0): x cp.arange(n, dtypecp.float32) y cp.ones(n, dtypecp.float32) wp.launch(saxpy, dimn, inputs[x, y, 1.0], devicecuda:0)注意直接传递 CUDA 数组时必须确保数组所在的设备与内核启动的设备一致如devicecuda:0对应cp.cuda.Device(0)否则会造成设备不匹配错误或隐式同步开销。这一约束同样适用于后面介绍的所有 CUDA 转换路径。这种方式的便利性体现在无需调用转换函数主要限制是标准数组接口不携带梯度信息因此只适合不涉及自动微分的算法。NumPy 互操作从 Warp 数组到 NumPyWarp 数组通过array.numpy()方法转换为 NumPy 数组。当 Warp 数组位于cpu设备时该方法返回零拷贝视图直接指向 Warp 底层分配当数组位于cuda设备时会先复制回临时缓冲区再拷贝给 NumPyw wp.array([1.0, 2.0, 3.0], dtypefloat, devicecpu) a np.array(w) # 通过 __array_interface__ 构造 print(a) # [1. 2. 3.]Warp CPU 数组实现了__array_interface__协议因此可以直接用np.array(w)构造 NumPy 数组无需显式转换。数据类型映射工具Warp 提供了方便的数据类型转换工具用于在两种类型系统之间映射warp_type wp.float32 ... numpy_type wp.dtype_to_numpy(warp_type) ... a wp.zeros(n, dtypewarp_type) b np.zeros(n, dtypenumpy_type)wp.dtype_to_numpy的实现位于 warp/_src/types.py其内部通过warp_type_to_np_dtype查表完成映射对不支持的 Warp 类型会抛出TypeError。从 NumPy 到 Warp要基于 NumPy 数组创建 Warp 数组使用wp.from_numpy源码位于 warp/_src/context.py或将 NumPy 数组直接作为wp.array构造函数的data参数传入。from_numpy支持dtype、shape、device、requires_grad、retain_grad等参数用于控制目标类型、放置设备与梯度行为。CuPy / Numba 互操作Warp GPU 数组实现了__cuda_array_interface__协议因此可以与其他 Python GPU 框架直接共享数据。这意味着CuPy、Numba 可以直接使用 Warp GPU 数组作为输入Warp 数组可以从任何暴露__cuda_array_interface__的对象创建这类对象也可以不创建 Warp 数组对象直接传给 Warp 内核见上文直接传递数组。由于该协议是纯内存描述指针、形状、步长、dtype不携带梯度信息所以这条路径没有梯度感知能力适合非求导场景。Paddle 互操作Warp 提供辅助函数在 Warp 数组与 Paddle 张量之间互转w wp.array([1.0, 2.0, 3.0], dtypefloat, devicecpu) # 转换为 Paddle 张量 t wp.to_paddle(w) # 从 Paddle 张量转换回来 w wp.from_paddle(t)这些辅助函数from_paddle位于 warp/_src/paddle.pyto_paddle位于同文件第 333 行不复制底层数据。与 PyTorch 路径一致梯度数组/张量会被转换为 Paddle autograd 张量使 Warp 数组可以参与 Paddle 的自动微分计算。Paddle 还提供 CUDA 流转换函数wp.stream_from_paddlewarp/_src/paddle.py用于把 Paddle CUDA 流转换为 Warp CUDA 流确保两个框架在共享零拷贝缓冲区时操作顺序正确。优化示例用wp.to_paddle声明优化变量当优化变量直接声明在 Warp 中时只需要一次wp.to_paddle调用即可把变量交给 Paddle 的 Adam 优化器——梯度由 Warp 的 tape 计算并写入 Warp 梯度缓冲区Paddle 优化器直接读取这些梯度import warp as wp import numpy as np import paddle wp.kernel() def loss(xs: wp.array2d[float], l: wp.array[float]): tid wp.tid() wp.atomic_add(l, 0, xs[tid, 0] ** 2.0 xs[tid, 1] ** 2.0) # 在 Warp 中初始化优化变量 xs wp.array(np.random.randn(100, 2), dtypewp.float32, requires_gradTrue) l wp.zeros(1, dtypewp.float32, requires_gradTrue) # 仅需一次 wp.to_paddle 调用Adam 使用 Warp 数组的梯度进行优化 opt paddle.optimizer.Adam(learning_rate0.1, parameters[wp.to_paddle(xs)]) tape wp.Tape() with tape: wp.launch(loss, dimlen(xs), inputs[xs], outputs[l], devicexs.device) for i in range(500): tape.zero() tape.backward(lossl) opt.step() l.zero_() wp.launch(loss, dimlen(xs), inputs[xs], outputs[l], devicexs.device) print(f{i}\tloss: {l.numpy()[0]})性能提示Paddle 遵循与 PyTorch 相同的调优模式——包括return_ctype参数跳过创建wp.array对象、直接返回低层数组描述符、直接传递张量依赖标准数组接口以及重用而非反复转换的原则。详见下文 PyTorch 性能调优章节这些经验对 Paddle 同样适用。DLPack 互操作Warp 支持 Python Array API 标准 v2022.12 中纳入的 DLPack 协议。DLPack 是一套与框架无关的共享内存描述机制允许在不拷贝数据的前提下跨框架传递数组。从外部框架导入wp.from_dlpack将外部数组导入 Warp 的标准方式是wp.from_dlpack()源码位于 warp/_src/dlpack.pywarp_array wp.from_dlpack(external_array)外部数组可以是 PyTorch 张量、JAX 数组或任何与该版本 DLPack 协议兼容的数组类型。从源码看from_dlpack会优先调用源的__dlpack__/__dlpack_device__接口见 warp/_src/dlpack.py对 CUDA 数组Warp 要求生产者producer在数组所在设备的当前 Warp 流上执行同步保证后续 Warp 内核对该数组的访问顺序正确。因此在同一设备上直接使用该数组通常是安全的无需额外同步对 CPU 数组不做流同步对 CUDA Hostpinned memory数组则与当前 CUDA 设备流同步。导出到外部框架framework.from_dlpack将 Warp 数组导出到外部框架的标准方式是使用对方框架的from_dlpack()函数jax_array jax.dlpack.from_dlpack(warp_array) torch_tensor torch.utils.dlpack.from_dlpack(warp_array) paddle_tensor paddle.utils.dlpack.from_dlpack(warp_array)对 CUDA 数组这会把消费方框架的当前流与 Warp 在数组设备上的当前流同步因此即使该数组此前在 Warp 内核中使用过包装后也能安全地在消费方框架中直接使用。使用 PyCapsule 的显式方式to_dlpack另一种共享方式是通过 PyCapsule 显式传递 DLPack 句柄生产者框架提供to_dlpack()函数消费方用from_dlpack()接收。这种方式适用于不支持 v2022.12 标准的老版本框架warp_array1 wp.from_dlpack(jax_array) warp_array2 wp.from_dlpack(torch.utils.dlpack.to_dlpack(torch_tensor)) warp_array3 wp.from_dlpack(paddle.utils.dlpack.to_dlpack(paddle_tensor)) jax_array jax.dlpack.from_dlpack(wp.to_dlpack(warp_array)) torch_tensor torch.utils.dlpack.from_dlpack(wp.to_dlpack(warp_array)) paddle_tensor paddle.utils.dlpack.from_dlpack(wp.to_dlpack(warp_array))Warp 侧导出接口wp.to_dlpackwarp/_src/dlpack.py返回包含DLManagedTensor的 PyCapsule可零拷贝转换为其他数组类型。从源码看它对结构化数组Structdtype会直接报错而向量/矩阵 dtype 会被展平为带额外内部维度的标量类型描述warp/_src/dlpack.py。性能权衡PyCapsule 方式一般更快因为它跳过了流同步但需要自行保证操作的顺序正确性。适合以下场景外部框架使用同步的 CUDA 默认流Warp 与外部框架使用同一条 CUDA 流已有其他同步机制在起作用。何时用 DLPack何时用专用转换器当存在框架专用转换器wp.to_torch、wp.to_paddle等时通常应优先使用它们因为DLPack 不携带梯度信息。如果 autograd 需要在 Warp 与其他框架之间流动请使用直接转换器。DLPack 的价值在于没有专用转换器可用时如 JAX 的数组互转在内部即经由 DLPack或者双方都是 DLPack 原生的场景。PyTorch 深度互操作流、图捕获、autograd 与性能调优Warp 对 PyTorch 的完整支持详见 docs/user_guide/interoperability/pytorch.rst以下是核心要点与主文档的零拷贝转换原则一脉相承。数组、设备与 dtype 转换wp.from_torch/wp.to_torchwarp/_src/torch.py 与第 326 行在不复制数据的前提下互转 Warp 数组与 PyTorch 张量并尽量把梯度数组与 PyTorch autograd 张量互转。同时提供设备与 dtype 的映射函数import torch torch_device wp.device_to_torch(cpu) torch_dtype wp.dtype_to_torch(wp.float32) t torch.ones(3, devicetorch_device, dtypetorch_dtype) warp_dtype wp.dtype_from_torch(t.dtype) warp_device wp.device_from_torch(t.device)从from_torch的源码warp/_src/torch.py可以看到一个关键实现细节标量 Warp dtype 会保留 PyTorch 张量的形状与步长因此非连续的张量通常可以直接包装而向量/矩阵 dtype 会消费张量的尾部连续分量维度如wp.vec2对应形状(..., 2)若尾部步长不连续则会抛出RuntimeError。此外当requires_gradTrue但张量尚未分配梯度时from_torch会用 Warp 分配一个零填充的梯度并挂回t.grad第 288-293 行——这正是下文延迟梯度分配问题的根源。默认情况下wp.zeros等分配函数使用 Warp 的 CUDA 分配器若希望分配来自 PyTorch 的 CUDA 缓存分配器可参考pytorch-cuda-caching-allocator的最小自定义分配器示例。流转换与 CUDA Graph 捕获流转换函数wp.stream_from_torch/wp.stream_to_torchwarp/_src/torch.py用于在两种框架的 CUDA 流之间互转。阻塞/非阻塞语义会被保留PyTorch 的默认流是阻塞的torch.cuda.Stream()创建的非默认流是非阻塞的而 Warp 创建的流是阻塞的。非阻塞流的垃圾回收风险与规避策略详见nonblocking_streams文档。由于 PyTorch 与 Warp 操作必须运行在同一条 CUDA 流上才能合并捕获 CUDA Graph而 PyTorch 默认的同步默认流不适合图捕获因此捕获前必须创建新流。两种捕获方式方式一用 PyTorch 流捕获转换为 Warp 流import torch import warp as wp wp.kernel def scale(a: wp.array[float], s: float): tid wp.tid() a[tid] a[tid] * s n 1024 * 1024 torch_device wp.device_to_torch(cuda:0) # 创建非默认 PyTorch 流并转换为 Warp 流 torch_stream torch.cuda.Stream(devicetorch_device) warp_stream wp.stream_from_torch(torch_stream) a wp.ones(n, dtypefloat, devicecuda:0) # 在共享流上捕获图 with wp.ScopedStream(warp_stream): with wp.ScopedCapture() as capture: wp.launch(scale, dimn, inputs[a, 2.0]) # 回放图 wp.capture_launch(capture.graph, streamwarp_stream)方式二用 Warp 流捕获让 PyTorch 使用 Warp 流import torch import warp as wp wp.kernel def scale(a: wp.array[float], s: float): tid wp.tid() a[tid] a[tid] * s n 1024 * 1024 a wp.ones(n, dtypefloat, devicecuda:0) # 让 PyTorch 使用 Warp 流 torch_stream wp.stream_to_torch(cuda:0) # 用 Warp 流捕获图 with wp.ScopedDevice(cuda:0), torch.cuda.stream(torch_stream): with wp.ScopedCapture() as capture: wp.launch(scale, dimn, inputs[a, 2.0]) # 回放图 wp.capture_launch(capture.graph)需要提醒的是许多 PyTorch 操作包含不可捕获的代码任意 PyTorch 代码的图捕获可能比较棘手可能需要进行预热warmup步骤。优化示例Warp 内核 PyTorch Adam与 Paddle 示例对称PyTorch 也有两个等价写法。wp.from_torch方向优化变量声明在 PyTorchWarp 通过零拷贝包装使用import warp as wp import torch wp.kernel() def loss(xs: wp.array2d[float], l: wp.array[float]): tid wp.tid() wp.atomic_add(l, 0, xs[tid, 0] ** 2.0 xs[tid, 1] ** 2.0) # requires_grad 使 Warp 能在 grad 缓冲区中累积梯度 xs torch.randn(100, 2, requires_gradTrue) l torch.zeros(1, requires_gradTrue) opt torch.optim.Adam([xs], lr0.1) wp_xs wp.from_torch(xs) wp_l wp.from_torch(l) tape wp.Tape() with tape: wp.launch(loss, dimlen(xs), inputs[wp_xs], outputs[wp_l], devicewp_xs.device) for i in range(500): tape.zero() tape.backward(losswp_l) # 计算梯度填充 xs.grad opt.step() # 更新 xs进而更新 wp_xs wp_l.zero_() wp.launch(loss, dimlen(xs), inputs[wp_xs], outputs[wp_l], devicewp_xs.device) print(f{i}\tloss: {l.item()})wp.to_torch方向优化变量声明在 Warp单次转换交给 PyTorch与上文 Paddle 示例结构完全相同只需把paddle.optimizer.Adam换成torch.optim.Adam([wp.to_torch(xs)], lr0.1)。Autograd 集成自定义算子与梯度缓冲区所有权将 Warp 内核插入 PyTorch 计算图有两条主流路径torch.autograd.FunctionPyTorch 2.3.1定义forward/backwardforward 中把入参张量映射为 Warp 数组后正常启动内核backward 中用wp.launch(..., adjointTrue)启动同一内核的伴随adjoint版本或依赖 Warp 的 tape。由于from_torch/to_torch是零拷贝转换backward 中收到的grad_output必须视为外部拥有的缓冲区PyTorch 可能在多次 backward 间复用同一张量绝不能把wp.from_torch(grad_output)直接赋给某个输出数组的.grad属性。梯度缓冲区所有权规则总结如下模式Warp 使用的缓冲区可安全复用/保留指引output.grad wp.from_torch(grad_output)外部 PyTorch 缓冲区否避免。Warp 可能消费或清零 PyTorch 打算复用的存储。tape.backward(grads{output: external_grad})且output.grad is Noneexternal_grad本身Tape 将其作为output.grad否先为output分配独立的梯度缓冲区。tape.backward(grads{output: external_grad})且output已拥有.grad已拥有的 Warp 缓冲区external_grad被拷贝进去是推荐用于外部 PyTorch 梯度。wp.to_torch(input.grad)Warp 梯度缓冲区的零拷贝视图仅到该缓冲区被修改前若 PyTorch 需保留梯度在tape.zero()前调用.clone()。此外若 backward 依赖 PyTorch 输入的 forward 值请用ctx.save_for_backward()保存原张量即使 Warp 包装的是 detached 视图ctx.saved_tensors的访问会让 PyTorch 在 Warp 读取共享存储前检测到就地修改CUDA 上应通过wp.stream_from_torchwp.ScopedStream让 Warp 工作运行在 PyTorch 活跃流上。文档中的完整 Rosenbrock 示例docs/user_guide/interoperability/pytorch.rst展示了forward/backward的完整实现。PyTorch 自定义算子PyTorch 2.4.0PyTorch 2.4 引入的 custom operators 把任意 Python 函数包括 Warp 调用视为不透明可调用对象阻止torch.compile()追踪进入从而让包含 Warp 内核启动的 forward 图可以被torch.compile()安全加速。其模式为用torch.library.custom_op注册 forward 与 backward 算子、用register_fake提供元数据形状、用register_autograd挂接 backward随后即可把整个 forward 包进torch.compile(fullgraphTrue)的函数中。性能调优return_ctype、直接传张量与转换复用wp.from_torch虽然不拷贝数据但每次转换仍有 CPU 开销创建wp.array对象。高频转换会拖累整体性能调优三板斧重用已转换的数组反复from_torch同一张量应避免。一次性转换后循环内直接复用x_t torch.arange(n, dtypetorch.float32, devicedevice) y_t torch.ones(n, dtypetorch.float32, devicedevice) x_w wp.from_torch(x_t) y_w wp.from_torch(y_t) for i in range(10): wp.launch(saxpy, dimn, inputs[x_w, y_w, 1.0], devicedevice)return_ctypeTrue当无法复用每轮迭代都构造新张量时wp.from_torch(x_t, return_ctypeTrue)跳过wp.array对象构造直接返回低层数组描述符C 结构可传给 Warp 内核但不能用于其他需要wp.array的地方for n in range(1, 10): x_t torch.arange(n, dtypetorch.float32, devicedevice) y_t torch.ones(n, dtypetorch.float32, devicedevice) x_ctype wp.from_torch(x_t, return_ctypeTrue) y_ctype wp.from_torch(y_t, return_ctypeTrue) wp.launch(saxpy, dimn, inputs[x_ctype, y_ctype, 1.0], devicedevice)直接传张量把 PyTorch 张量直接传给 Warp 内核依赖__cuda_array_interface__完全省去转换函数代价是不处理梯度适合无求导算法。仓库提供了可运行基准 warp/examples/benchmarks/benchmark_interop_torch.py用以下命令对比三种模式python -m warp.examples.benchmarks.benchmark_interop_torch文档中的样本输出显示from_torch(...)最慢约 5095 msfrom_torch(..., return_ctypeTrue)最快约 2113 ms直接传张量居中约 2950 ms——直接传张量虽省去了临时 Warp 数组但访问 PyTorch 张量的__cuda_array_interface__属性有按需初始化的开销。若在这些模式之上构建缓存例如以张量data_ptr()或 Warp 数组描述符为键请在底层 Warp 数组释放时失效缓存——新分配可能复用同一内存地址但尺寸/形状/dtype 不同指针相等性不能作为安全缓存键。案例研究PyTorch 延迟梯度分配导致的同步开销PyTorch 对梯度张量采用延迟分配策略requires_gradTrue时并不会立即分配梯度内存而是在 backward 过程中按需分配。问题在于wp.from_torch遇到有requires_gradTrue但没有分配梯度的张量时会强制立即分配梯度见 warp/_src/torch.py当 PyTorch 随后发现外部框架已分配其梯度张量时必须执行昂贵的设备级同步来保证正确性。若每轮迭代都用.clone().detach().requires_grad_(True)新建张量则该惩罚每轮都会发生。如上图NVIDIA Nsight Systems 时间线所示Warp 内核启动与 PyTorch 操作之间出现明显的设备级同步间隙。仓库文档中的案例以 N3 亿元素负载测得的对比为Baseline (with synchronization overhead): 98.02 ms Solution A (requires_gradFalse): 22.59 ms (4.3x faster) Solution B (detach): 22.11 ms (4.4x faster) Solution C (pre-allocate): 28.62 ms (3.4x faster)三种解决方案各有适用场景Solution Awp.from_torch(..., requires_gradFalse)——最简单禁止 Warp 自动分配梯度适合手动管理 forward/梯度张量的场景Solution Bdetach 张量——用x.detach()把张量移出 PyTorch 计算图并清除requires_grad明确梯度管理在 PyTorch autograd 之外Solution C用 PyTorch 分配器预分配梯度——a.grad torch.empty_like(a)分析梯度内核路径或用专用零填充缓冲区 wp.from_torch(..., gradctx.grad_a)显式挂接Warp tape 路径。tape 路径中输出数组也应通过gradctx.grad_output获得自有梯度缓冲区使tape.backward(grads{...})把外部梯度拷贝进自有存储而非被 tape 收养backward 中在tape.zero()之前先.clone()出要返回给 PyTorch 的梯度。Solution C 是使用 Warp tape 或需要访问.grad时的必需方案。若你的工作负载在整个迭代中复用同一批张量梯度已分配则不会出现延迟或非延迟的梯度分配也就没有同步开销。JAX 深度互操作FFI 内核、vmap、自动微分与分布式JAX 的互操作支持详见 docs/user_guide/interoperability/jax.rst。JAX 数组互转内部使用 DLPack 协议零拷贝交换数据warp_array wp.from_jax(jax_array) jax_array wp.to_jax(warp_array)实现见 warp/_src/jax/init.py。追求更优性能与流同步控制时也可直接用 DLPack 协议。把 Warp 内核作为 JAX 原语jax_kerneljax_kernel源码位于 warp/_src/jax/ffi.py把单个 Warp 内核包装成 JAX 原语可在 jitted JAX 函数内调用import warp as wp import jax import jax.numpy as jnp from warp import jax_kernel wp.kernel def triple_kernel(input: wp.array[float], output: wp.array[float]): tid wp.tid() output[tid] 3.0 * input[tid] # 从 Warp 内核创建 JAX 原语 jax_triple jax_kernel(triple_kernel) jax.jit def f(): x jnp.arange(0, 64, dtypejnp.float32) return jax_triple(x) print(f())设备选择JAX 依据调用被 lower 到的设备选择 FFI 实现同一包装器在 CPU 与 CUDA 上均可工作无需 Warp 设备参数——Warp 直接包装 XLA 缓冲区、不拷贝数据。同一 jitted 函数可分别对jax.devices(cpu)[0]与jax.devices(cuda)[0]上的输入运行。输入输出语义内核定义中输入参数必须位于输出参数之前至少需要一个输出数组允许无输入内核输出个数用num_outputs指定默认 1标量输入必须为 JAX 中的常量或静态值traced 标量会抛异常可用partial(jax.jit, static_argnames[s])使标量静态化默认按第一个输入数组形状推断 launch 维度需要时可用launch_dims覆盖输出数组形状默认由 launch 维度决定也可用output_dims自定义支持整数的 1D 形状、元组/列表的多维形状以及{b: n, c: m}形式的按输出字典无输入内核必须显式传launch_dims以确定输出形状向量/矩阵数组JAX 没有对应类型分量被打包为额外内部维度——wp.vec3数组对应 JAX 形状(..., 3)wp.mat22对应(..., 2, 2)。output_dims同时接受两种约定Warp 形状或 JAX 形状CUDA 上默认每块 256 线程可用block_dim调整构建包装器时固定tile 内核需把执行宽度作为尾随 launch 维度传入并保持output_dims为逻辑输出形状。VMAP 支持vmap_method参数默认broadcast_all控制回调在jax.vmap下的变换方式可在构建jax_kernel时设定默认值也可在单次调用时覆盖。对含 in-out 参数的内核用in_out_argnames[sums]声明vmap 中可指定in_axes匹配批量维度对 launch 维度不同于首个数组形状的内核如查表lookup_kernel用functools.partial(jax_lookup, launch_dims50)传递自定义参数——注意launch_dims/output_dims不应包含批量维度批量由 vmap 自动处理。自动微分实验性传enable_backwardTrue给jax_kernel即可为内核挂接自定义 VJP使jax.grad可对 Warp 内核求导forward 与伴随由同一launch_dims驱动。当前限制标量输入必须是 JAX 静态参数梯度仅针对可微分的数组输入返回静态标量不在梯度元组中in_out_argnames与output_dims在enable_backwardTrue时不支持launch_dims在enable_backwardTrue时于构建期固定不可逐调用覆盖当输入数组维度多于内核wp.tid()迭代空间如 LBM 分布(Q, nx, ny, nz)时务必显式传launch_dims空间维度否则伴随内核会通过atomic_add按外轴大小过度累积梯度。jax_callable多内核函数与图捕获jax_callablewarp/_src/jax/ffi.py允许从 JAX 调用会启动多个内核的 Python 函数目标函数需像 Warp 内核一样带参数类型注解。其输入输出语义与jax_kernel类似差异在于不接受launch_dims由目标函数自行启动内核接受graph_mode参数控制 CUDA 图捕获方式——JAX默认让 JAX 捕获可作为子图、WARPWarp 捕获输入输出缓冲区地址匹配时复用捕获图、WARP_STAGED/WARP_STAGED_EX对稳定 staging 缓冲区捕获拷贝作为图节点/图外提交、NONE禁用用于含主机同步等不可捕获操作。staged 模式会占用额外显存可用stage_in_argnames/stage_out_argnames限制逐调用拷贝范围并用graph_cache_max限制缓存图数量、clear_jax_callable_graph_cache()释放缓存。module_preload_modeCURRENT_DEVICE/ALL_DEVICES/NONE控制模块预加载范围。完整示例见 warp/examples/interop/example_jax_callable.py 与 warp/examples/interop/example_jax_kernel.py。分布式计算shard_mapWarp 可与 JAX 的shard_map结合实现多 GPU 分布式计算。程序开头必须先jax.distributed.initialize()在任何其他 JAX 操作之前。在shard_map的 sharded 算子内部每个设备只处理本地分片Warp 内核作用于本地分片并返回同样形状的结果import warp as wp import jax import jax.numpy as jnp from jax.sharding import PartitionSpec as P from jax.experimental.multihost_utils import process_allgather as allgather from jax.experimental.shard_map import shard_map from warp import jax_kernel import numpy as np jax.distributed.initialize() num_gpus jax.device_count() wp.kernel def multiply_by_two_kernel(a_in: wp.array[float], a_out: wp.array[float]): index wp.tid() a_out[index] a_in[index] * 2.0 jax_warp_multiply jax_kernel(multiply_by_two_kernel) def warp_distributed_operator(a_in): def _sharded_operator(a_in): # 每个设备上 a_in 是本地分片形状 (M/N,) result warp_multiply(a_in)[0] return result return shard_map( _sharded_operator, meshjax.sharding.Mesh(np.array(jax.devices()), x), in_specs(P(x),), # 输入沿 x 轴分片 out_specsP(x), # 输出同样沿 x 轴分片 check_repFalse, )(a_in)运行多 GPU 程序需要安装 Open MPI并用mpirun启动mpirun -np NUM_OF_GPUS python filename.py结语如何选择合适的互操作路径综合本指南选择路径的核心判断依据是是否需要梯度以及性能敏感度最快速接入任意实现了__array_interface__/__cuda_array_interface__的数组直接传入wp.launch零转换开销但无梯度需要 autograd优先使用框架专用转换器wp.from_torch/wp.to_torch、wp.from_paddle/wp.to_paddle它们零拷贝且携带梯度注意高频转换时的return_ctypeTrue、复用与直接传张量三种调优手段以及 PyTorch 延迟梯度分配陷阱无专用转换器或双方 DLPack 原生使用wp.from_dlpack/ 框架from_dlpack需理解流同步语义追求极致性能且同步可由其他机制保证时用 PyCapsule 的to_dlpack路径JAX 生态数组互转走from_jax/to_jax内部即 DLPack内核级集成用jax_kernel/jax_callable配合 vmap、autodiff实验性与shard_map覆盖从单卡到分布式的完整需求。以上接口的完整 API 说明可进一步查阅 docs/api_reference/warp.rst相关转换函数与 FFI 实现的源码集中在 warp/_src/torch.py、warp/_src/paddle.py、warp/_src/dlpack.py 与 warp/_src/jax/ 目录下。【免费下载链接】warpA Python framework for GPU-accelerated simulation, robotics, and machine learning.项目地址: https://gitcode.com/GitHub_Trending/warp/warp创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考