news 2026/9/17 3:11:12

Warp 互操作性实战指南:在 NumPy、PyTorch、JAX、Paddle 与 DLPack 之间零拷贝共享 GPU 数据

作者头像

张小明

前端开发工程师

1.2k 24
文章封面图
Warp 互操作性实战指南:在 NumPy、PyTorch、JAX、Paddle 与 DLPack 之间零拷贝共享 GPU 数据

Warp 互操作性实战指南:在 NumPy、PyTorch、JAX、Paddle 与 DLPack 之间零拷贝共享 GPU 数据

【免费下载链接】warpA Python framework for GPU-accelerated simulation, robotics, and machine learning.项目地址: https://gitcode.com/GitHub_Trending/warp/warp

Warp(warp,A 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.12),JAX、PyTorch、Paddle 均支持。

从源码结构看,协议支持被集中封装在 warp/_src/context.py(from_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_kernel
Paddlewp.from_paddle/wp.to_paddle
CuPy / Numba__cuda_array_interface__协议
DLPackwp.from_dlpack/framework.from_dlpack

上表中"梯度感知"为Yes的框架,其转换函数会在可用时把 Warp 的梯度数组与对应框架的 autograd 张量互转,从而让 Warp 数组参与对方框架的反向传播计算。JAX 的梯度感知能力通过jax_kernel(FFI)包装器提供,而非数组转换本身。

直接传递数组:最快的上手路径

任何实现了__array_interface__(CPU)或__cuda_array_interface__(GPU)的对象,都可以不调用任何转换函数直接传入wp.launchinputs。这是大多数场景下最快的接入方式——省去了创建 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, dtype=np.float32) y = np.ones(n, dtype=np.float32) wp.launch(saxpy, dim=n, inputs=[x, y, 1.0], device="cpu")

CUDA 端,同样的模式适用于 CuPy、PyTorch 或任何暴露__cuda_array_interface__的框架:

import cupy as cp with cp.cuda.Device(0): x = cp.arange(n, dtype=cp.float32) y = cp.ones(n, dtype=cp.float32) wp.launch(saxpy, dim=n, inputs=[x, y, 1.0], device="cuda:0")

注意:直接传递 CUDA 数组时,必须确保数组所在的设备与内核启动的设备一致(如device="cuda:0"对应cp.cuda.Device(0)),否则会造成设备不匹配错误或隐式同步开销。这一约束同样适用于后面介绍的所有 CUDA 转换路径。

这种方式的便利性体现在无需调用转换函数;主要限制是标准数组接口不携带梯度信息,因此只适合不涉及自动微分的算法。

NumPy 互操作

从 Warp 数组到 NumPy

Warp 数组通过array.numpy()方法转换为 NumPy 数组。当 Warp 数组位于cpu设备时,该方法返回零拷贝视图,直接指向 Warp 底层分配;当数组位于cuda设备时,会先复制回临时缓冲区再拷贝给 NumPy:

w = wp.array([1.0, 2.0, 3.0], dtype=float, device="cpu") 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, dtype=warp_type) b = np.zeros(n, dtype=numpy_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支持dtypeshapedevicerequires_gradretain_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], dtype=float, device="cpu") # 转换为 Paddle 张量 t = wp.to_paddle(w) # 从 Paddle 张量转换回来 w = wp.from_paddle(t)

这些辅助函数(from_paddle位于 warp/_src/paddle.py,to_paddle位于同文件第 333 行)不复制底层数据。与 PyTorch 路径一致,梯度数组/张量会被转换为 Paddle autograd 张量,使 Warp 数组可以参与 Paddle 的自动微分计算。

Paddle 还提供 CUDA 流转换函数wp.stream_from_paddle(warp/_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), dtype=wp.float32, requires_grad=True) l = wp.zeros(1, dtype=wp.float32, requires_grad=True) # 仅需一次 wp.to_paddle 调用,Adam 使用 Warp 数组的梯度进行优化 opt = paddle.optimizer.Adam(learning_rate=0.1, parameters=[wp.to_paddle(xs)]) tape = wp.Tape() with tape: wp.launch(loss, dim=len(xs), inputs=[xs], outputs=[l], device=xs.device) for i in range(500): tape.zero() tape.backward(loss=l) opt.step() l.zero_() wp.launch(loss, dim=len(xs), inputs=[xs], outputs=[l], device=xs.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.py):

warp_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 Host(pinned 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_dlpack(warp/_src/dlpack.py)返回包含DLManagedTensor的 PyCapsule,可零拷贝转换为其他数组类型。从源码看,它对结构化数组(Structdtype)会直接报错,而向量/矩阵 dtype 会被展平为带额外内部维度的标量类型描述(warp/_src/dlpack.py)。

性能权衡:PyCapsule 方式一般更快,因为它跳过了流同步,但需要自行保证操作的顺序正确性。适合以下场景:

  • 外部框架使用同步的 CUDA 默认流;
  • Warp 与外部框架使用同一条 CUDA 流;
  • 已有其他同步机制在起作用。

何时用 DLPack,何时用专用转换器

当存在框架专用转换器(wp.to_torchwp.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_torch(warp/_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, device=torch_device, dtype=torch_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_grad=True但张量尚未分配梯度时,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_torch(warp/_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(device=torch_device) warp_stream = wp.stream_from_torch(torch_stream) a = wp.ones(n, dtype=float, device="cuda:0") # 在共享流上捕获图 with wp.ScopedStream(warp_stream): with wp.ScopedCapture() as capture: wp.launch(scale, dim=n, inputs=[a, 2.0]) # 回放图 wp.capture_launch(capture.graph, stream=warp_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, dtype=float, device="cuda: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, dim=n, inputs=[a, 2.0]) # 回放图 wp.capture_launch(capture.graph)

需要提醒的是:许多 PyTorch 操作包含不可捕获的代码,任意 PyTorch 代码的图捕获可能比较棘手,可能需要进行预热(warmup)步骤。

优化示例:Warp 内核 + PyTorch Adam

与 Paddle 示例对称,PyTorch 也有两个等价写法。wp.from_torch方向(优化变量声明在 PyTorch,Warp 通过零拷贝包装使用):

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_grad=True) l = torch.zeros(1, requires_grad=True) opt = torch.optim.Adam([xs], lr=0.1) wp_xs = wp.from_torch(xs) wp_l = wp.from_torch(l) tape = wp.Tape() with tape: wp.launch(loss, dim=len(xs), inputs=[wp_xs], outputs=[wp_l], device=wp_xs.device) for i in range(500): tape.zero() tape.backward(loss=wp_l) # 计算梯度,填充 xs.grad opt.step() # 更新 xs(进而更新 wp_xs) wp_l.zero_() wp.launch(loss, dim=len(xs), inputs=[wp_xs], outputs=[wp_l], device=wp_xs.device) print(f"{i}\tloss: {l.item()}")

wp.to_torch方向(优化变量声明在 Warp,单次转换交给 PyTorch)与上文 Paddle 示例结构完全相同,只需把paddle.optimizer.Adam换成torch.optim.Adam([wp.to_torch(xs)], lr=0.1)

Autograd 集成:自定义算子与梯度缓冲区所有权

将 Warp 内核插入 PyTorch 计算图有两条主流路径:

  1. torch.autograd.Function(PyTorch <= 2.3.1):定义forward/backward,forward 中把入参张量映射为 Warp 数组后正常启动内核;backward 中用wp.launch(..., adjoint=True)启动同一内核的伴随(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_torch+wp.ScopedStream让 Warp 工作运行在 PyTorch 活跃流上。文档中的完整 Rosenbrock 示例(docs/user_guide/interoperability/pytorch.rst)展示了forward/backward的完整实现。

  1. PyTorch 自定义算子(PyTorch >= 2.4.0):PyTorch 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(fullgraph=True)的函数中。

性能调优:return_ctype、直接传张量与转换复用

wp.from_torch虽然不拷贝数据,但每次转换仍有 CPU 开销(创建wp.array对象)。高频转换会拖累整体性能,调优三板斧:

  1. 重用已转换的数组:反复from_torch同一张量应避免。一次性转换后循环内直接复用:
x_t = torch.arange(n, dtype=torch.float32, device=device) y_t = torch.ones(n, dtype=torch.float32, device=device) x_w = wp.from_torch(x_t) y_w = wp.from_torch(y_t) for i in range(10): wp.launch(saxpy, dim=n, inputs=[x_w, y_w, 1.0], device=device)
  1. return_ctype=True:当无法复用(每轮迭代都构造新张量)时,wp.from_torch(x_t, return_ctype=True)跳过wp.array对象构造,直接返回低层数组描述符(C 结构),可传给 Warp 内核但不能用于其他需要wp.array的地方:
for n in range(1, 10): x_t = torch.arange(n, dtype=torch.float32, device=device) y_t = torch.ones(n, dtype=torch.float32, device=device) x_ctype = wp.from_torch(x_t, return_ctype=True) y_ctype = wp.from_torch(y_t, return_ctype=True) wp.launch(saxpy, dim=n, inputs=[x_ctype, y_ctype, 1.0], device=device)
  1. 直接传张量:把 PyTorch 张量直接传给 Warp 内核(依赖__cuda_array_interface__),完全省去转换函数;代价是不处理梯度,适合无求导算法。

仓库提供了可运行基准 warp/examples/benchmarks/benchmark_interop_torch.py,用以下命令对比三种模式:

python -m warp.examples.benchmarks.benchmark_interop_torch

文档中的样本输出显示:from_torch(...)最慢(约 5095 ms),from_torch(..., return_ctype=True)最快(约 2113 ms),直接传张量居中(约 2950 ms)——直接传张量虽省去了临时 Warp 数组,但访问 PyTorch 张量的__cuda_array_interface__属性有按需初始化的开销。若在这些模式之上构建缓存(例如以张量data_ptr()或 Warp 数组描述符为键),请在底层 Warp 数组释放时失效缓存——新分配可能复用同一内存地址但尺寸/形状/dtype 不同,指针相等性不能作为安全缓存键。

案例研究:PyTorch 延迟梯度分配导致的同步开销

PyTorch 对梯度张量采用延迟分配策略:requires_grad=True时并不会立即分配梯度内存,而是在 backward 过程中按需分配。问题在于:wp.from_torch遇到有requires_grad=True但没有分配梯度的张量时,会强制立即分配梯度(见 warp/_src/torch.py);当 PyTorch 随后发现外部框架已分配其梯度张量时,必须执行昂贵的设备级同步来保证正确性。若每轮迭代都用.clone().detach().requires_grad_(True)新建张量,则该惩罚每轮都会发生

如上图(NVIDIA Nsight Systems 时间线)所示,Warp 内核启动与 PyTorch 操作之间出现明显的设备级同步间隙。仓库文档中的案例以 N=3 亿元素负载测得的对比为:

Baseline (with synchronization overhead): 98.02 ms Solution A (requires_grad=False): 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 A:wp.from_torch(..., requires_grad=False)——最简单,禁止 Warp 自动分配梯度,适合手动管理 forward/梯度张量的场景;
  • Solution B:detach 张量——用x.detach()把张量移出 PyTorch 计算图并清除requires_grad,明确"梯度管理在 PyTorch autograd 之外";
  • Solution C:用 PyTorch 分配器预分配梯度——a.grad = torch.empty_like(a)(分析梯度内核路径),或用专用零填充缓冲区 +wp.from_torch(..., grad=ctx.grad_a)显式挂接(Warp tape 路径)。tape 路径中输出数组也应通过grad=ctx.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_kernel

jax_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, dtype=jnp.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_dims=50)传递自定义参数——注意launch_dims/output_dims不应包含批量维度,批量由 vmap 自动处理。

自动微分(实验性)

enable_backward=Truejax_kernel即可为内核挂接自定义 VJP,使jax.grad可对 Warp 内核求导(forward 与伴随由同一launch_dims驱动)。当前限制:

  • 标量输入必须是 JAX 静态参数;
  • 梯度仅针对可微分的数组输入返回(静态标量不在梯度元组中);
  • in_out_argnamesoutput_dimsenable_backward=True时不支持;
  • launch_dimsenable_backward=True时于构建期固定,不可逐调用覆盖
  • 当输入数组维度多于内核wp.tid()迭代空间(如 LBM 分布(Q, nx, ny, nz))时,务必显式传launch_dims(空间维度),否则伴随内核会通过atomic_add按外轴大小过度累积梯度。

jax_callable:多内核函数与图捕获

jax_callable(warp/_src/jax/ffi.py)允许从 JAX 调用会启动多个内核的 Python 函数,目标函数需像 Warp 内核一样带参数类型注解。其输入输出语义与jax_kernel类似,差异在于:不接受launch_dims(由目标函数自行启动内核);接受graph_mode参数控制 CUDA 图捕获方式——JAX(默认,让 JAX 捕获,可作为子图)、WARP(Warp 捕获,输入输出缓冲区地址匹配时复用捕获图)、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_map

Warp 可与 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, mesh=jax.sharding.Mesh(np.array(jax.devices()), "x"), in_specs=(P("x"),), # 输入沿 'x' 轴分片 out_specs=P("x"), # 输出同样沿 'x' 轴分片 check_rep=False, )(a_in)

运行多 GPU 程序需要安装 Open MPI,并用mpirun启动:

mpirun -np <NUM_OF_GPUS> python <filename>.py

结语:如何选择合适的互操作路径

综合本指南,选择路径的核心判断依据是是否需要梯度以及性能敏感度

  1. 最快速接入:任意实现了__array_interface__/__cuda_array_interface__的数组直接传入wp.launch,零转换开销,但无梯度;
  2. 需要 autograd:优先使用框架专用转换器(wp.from_torch/wp.to_torchwp.from_paddle/wp.to_paddle),它们零拷贝且携带梯度;注意高频转换时的return_ctype=True、复用与直接传张量三种调优手段,以及 PyTorch 延迟梯度分配陷阱;
  3. 无专用转换器或双方 DLPack 原生:使用wp.from_dlpack/ 框架from_dlpack,需理解流同步语义;追求极致性能且同步可由其他机制保证时,用 PyCapsule 的to_dlpack路径;
  4. 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),仅供参考

版权声明: 本文来自互联网用户投稿,该文观点仅代表作者本人,不代表本站立场。本站仅提供信息存储空间服务,不拥有所有权,不承担相关法律责任。如若内容造成侵权/违法违规/事实不符,请联系邮箱:809451989@qq.com进行投诉反馈,一经查实,立即删除!
网站建设 2026/9/17 3:10:02

anarlog windows 插件权限体系解析:Tauri 命令 ACL 参考与实践指南

anarlog windows 插件权限体系解析&#xff1a;Tauri 命令 ACL 参考与实践指南 【免费下载链接】anarlog Open source Granola AI Alternative 项目地址: https://gitcode.com/GitHub_Trending/hy/anarlog 本篇技术指南围绕 anarlog 桌面端 windows&#xff08;tauri-pl…

作者头像 李华
网站建设 2026/9/17 3:09:39

CANoe CAPL实战:8个车载网络高频场景的工程化解决方案

/* MD / 富文本中的 .toc(含博客园搬家等嵌套结构);.toc-box 在侧栏,不受影响 */#content_views .toc,/* 编辑器常在目录前后插入空 p(:empty 仍占 20px),一并去掉避免顶空隙 */#content_views.markdown_views > p:empty:has(+ .toc),#content_views.markdown_views …

作者头像 李华
网站建设 2026/9/17 3:09:36

AI短视频自动化工作流:Sora+CapCut全链路实操指南

/* MD / 富文本中的 .toc(含博客园搬家等嵌套结构);.toc-box 在侧栏,不受影响 */#content_views .toc,/* 编辑器常在目录前后插入空 p(:empty 仍占 20px),一并去掉避免顶空隙 */#content_views.markdown_views > p:empty:has(+ .toc),#content_views.markdown_views …

作者头像 李华
网站建设 2026/9/17 3:08:21

RailWay容器托管平台部署实践:边界、环境变量与故障排查

上个月我把一个断断续续跑了两年多的小后端从自己手动维护的环境里搬到了 RailWay 这个免费容器托管平台上&#xff0c;搬完当天晚上我就把之前写好的一堆定时重启脚本、日志切割脚本和证书续期脚本全删了。说这话不是劝所有人都去搬&#xff0c;而是想聊聊当一个「容器托管平台…

作者头像 李华
网站建设 2026/9/17 3:07:55

算法题中的指针类型题目:核心考点、解题套路与常见误区

最近在刷算法题的朋友应该有感受&#xff0c;链表、二叉树的题目十道里有七八道都在折腾指针。尤其是C/C选手&#xff0c;写双指针、快慢指针时经常被一两个星号搞得晕头转向——改了指针本身还是改指针指向的内容&#xff1f;改完下一个节点该接谁&#xff1f;一旦想不清楚&am…

作者头像 李华
网站建设 2026/9/17 3:07:30

Android美颜相机实现:CameraX取帧、美颜算法与GLSL渲染管线

简介&#xff1a;面向安卓开发与毕业设计人群的这份项目资料&#xff0c;围绕实现一款类似美颜相机、美图秀秀的应用展开&#xff0c;覆盖实时美颜、照片编辑、滤镜特效等核心功能需求。内容系统梳理了安卓开发基础、相机与相机第二代相机接口的调用、使用开放计算机视觉库进行…

作者头像 李华