尧图网站设计 尧图网站设计YAOTU DESIGN
ARTICLE DETAIL

资讯详情

深耕网站设计与一线实操的经验洞察。

如何用 DLPack 协议在 PyArrow 与张量框架之间交换数据

如何用 DLPack 协议在 PyArrow 与张量框架之间交换数据 如何用 DLPack 协议在 PyArrow 与张量框架之间交换数据【免费下载链接】arrowApache Arrow is the universal columnar format and multi-language toolbox for fast data interchange and in-memory analytics项目地址: https://gitcode.com/GitHub_Trending/arrow3/arrow如果你的数据管道一端是 PyArrow例如从 Parquet、CSV 或 Flight 读出的列式数据另一端是 NumPy、PyTorch 或 JAX 这类张量框架直接在两个体系之间复制内存既慢又浪费。DLPack 协议就是为此设计的它是一套稳定的内存中数据结构允许支持该协议的框架之间共享多维数组的内存而 PyArrow 在pa.Array、pa.Tensor和pa.FixedShapeTensorArray上实现了这一协议。本文的目标是让你能按文档给出的一条连续路径把 CPU 上的 PyArrow 数值数组零拷贝地交给张量框架或把张量框架里的数组零拷贝地导入 PyArrow并知道在哪些边界上协议会拒绝数据。准备环境PyArrow 的安装方式见 安装文档pip install pyarrow按文档说明PyArrow 兼容 Python 3.11、3.12、3.13 和 3.14且定期在 Windows、macOS 和主流 Linux 发行版上构建和测试建议使用 64 位系统。NumPy 是 PyArrow 的可选依赖文档要求 NumPy 2.0 或更高本文示例用 NumPy 作为 DLPack 消费方PyTorch 和 JAX 的示例则以各自框架的安装方式为准。开始操作前先明确 PyArrow 侧的实现范围这决定了你能直接导出的对象类型PyArrow 对象DLPack 行为pa.Tensor可产出和消费任意 shape、strides 的通用 DLPack tensorpa.Array被刻意限制为只能产出和消费一维、连续contiguous的 tensorpa.FixedShapeTensorArray可产出和消费“最外层维度拥有最大 stride”的 tensor该维度对应数组长度类型限制同样明确pa.Tensor和pa.Array只支持数值类型——整数、无符号整数和浮点数。另外PyArrow 当前对协议的实现只支持 CPU 设备上的数据。第一步把 PyArrow 数据导出给张量框架导出的语法是消费方调用from_dlpack(x, /, *, deviceNone, copyNone)它要求对象实现__dlpack__方法产出时共享内存而不复制数据。最常用的是把一维pa.Array直接交给 NumPy以下输出为 文档示例import pyarrow as pa import numpy as np array pa.array([2, 0, 2, 4]) np.from_dlpack(array) # 文档示例输出: array([2, 0, 2, 4])PyTorch 和 JAX 的用法同构文档示例import torch torch.from_dlpack(array) # 文档示例输出: tensor([2, 0, 2, 4])import jax jax.numpy.from_dlpack(array) # 文档示例输出: Array([2, 0, 2, 4], dtypeint32)如果你的数据是带张量内存布局的数组——比如数值类型的 fixed size list其内存表示与行主序row majortensor 相同——不能直接走 DLPack需要先显式转成pa.Tensor由它导出完整的多维 shapelist_array pa.array([[1.0, 2.0], [3.0, 4.0], [5.0, 6.0]], pa.list_(pa.float64(), 2)) tensor list_array.to_tensor() # 文档示例输出: # pyarrow.Tensor # type: double # shape: (3, 2) # strides: (16, 8) np.from_dlpack(tensor) # 文档示例输出: array([[1., 2.], [3., 4.], [5., 6.]])pa.FixedShapeTensorArray则可以直接导出数组长度成为最外层维度其后是元素 tensor 的 shapenested pa.array([[[1, 2], [3, 4]], [[5, 6], [7, 8]]], pa.list_(pa.list_(pa.int32(), 2), 2)) tensor_array pa.FixedShapeTensorArray.from_tensor(nested.to_tensor()) np.from_dlpack(tensor_array).shape # 文档示例输出: (2, 2, 2)第二步含 null 的数组如何处理DLPack 在含 null 的数组上会直接失败报错如下文档示例array_with_nulls pa.array([2, None, 4], pa.int32()) np.from_dlpack(array_with_nulls) # pyarrow.lib.ArrowTypeError: Can only use DLPack on arrays with no nulls.文档给出的绕过方式是用to_tensor(allow_nullsTrue)显式转成 tensor 再导出此时 null 位置持有未指定的值np.from_dlpack(array_with_nulls.to_tensor(allow_nullsTrue)) # 文档示例输出: array([2, ..., 4], dtypeint32)文档同时对比了另一种处理方式pa.compute.fill_null会显式修改数组数据来替换 null 值而allow_nullsTrue的转换是零开销的。如果下游会依赖 null 位置的值应选择fill_null而不是allow_nullsTrue。第三步把张量框架的数据导入 PyArrow反方向上任何实现 DLPack 协议的对象都可以被导入且不复制数据pa.Array.from_dlpack(np.array([2, 0, 2, 4])) # 文档示例输出: # pyarrow.lib.Int64Array object at ... # [2, 0, 2, 4] pa.Tensor.from_dlpack(np.array([[2, 0], [2, 4]], np.int32)) # 文档示例输出: # pyarrow.Tensor # type: int32 # shape: (2, 2) # strides: (8, 4)注意pa.Array.from_dlpack只接受一维连续的 tensor多维数据要导入为pa.FixedShapeTensorArray最外层维度会成为数组长度文档示例array pa.FixedShapeTensorArray.from_dlpack( np.arange(12, dtypenp.int32).reshape(3, 2, 2) ) array.type # FixedShapeTensorType(extensionarrow.fixed_shape_tensor[value_typeint32, shape[2,2], permutation[0,1]]) len(array) # 3两个方向都强调“共享内存”from_dlpack是“创建新对象但共享内存”。也就是说导入得到的 PyArrow 对象与源数组读写的是同一块内存修改一方会影响另一方这正是零拷贝的代价与收益所在验证时不要把它误判成数据漂移。验证方式与边界判断完成一次交换后可以从三个层面核对结果数值与形状把导入/导出的对象转回可打印形式如to_pylist()、shape、dtype与源数据比对。上文的文档示例输出可作为对照基准但它们只是示例不要当成固定预期值去断言其他输入。设备协议包含__dlpack_device__方法用于查询对象所在设备。PyArrow 的 DLPack 实现只支持 CPU 数据把 GPU 上的数组如pyarrow.cuda产生的CudaBuffer传给np.from_dlpack会抛出NotImplementedError测试文件 中对应的报错信息是DLPack support is implemented only for buffers on CPU device.类型支持非数值类型会被拒绝。同样的测试文件记录了具体报错例如空数组或普通 list 数组报DataType is not supported by DLPack spec一类的TypeErrorDataType is not compatible with DLPack specbit-packed 布尔值报Bit-packed boolean data type not supported by DLPack.参考文档协议说明与全部示例The DLPack Protocol安装与 Python 版本要求Installing PyArrow行为测试报错文案、零拷贝语义、设备检查python/pyarrow/tests/test_dlpack.py需要为库作者自己实现导出方时文档给出的接口是__dlpack__(self, *, streamNone, max_versionNone, dl_deviceNone, copyNone)与__dlpack_device__前者产出携带 DLPack 结构的 PyCapsule供对端的from_dlpack(x)调用。【免费下载链接】arrowApache Arrow is the universal columnar format and multi-language toolbox for fast data interchange and in-memory analytics项目地址: https://gitcode.com/GitHub_Trending/arrow3/arrow创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考
返回列表