NumPy 用户自定义 DType 的 Partition / Argpartition 支持:基于 ArrayMethod API 实现自定义分区算法
【免费下载链接】numpyThe fundamental package for scientific computing with Python.项目地址: https://gitcode.com/gh_mirrors/nu/numpy
本篇技术指南以 NumPy 官方发布说明(doc/release/upcoming_changes/31614.new_feature.rst)为切入点,系统讲解用户自定义 DType 如何通过 ArrayMethod API 为numpy.partition与numpy.argpartition注册自定义实现,覆盖注册入口、Spec 契约、参数结构与底层调用链。读完本文,你将掌握如何为自己的自定义 DType 编写可被 NumPy 分区操作调用的原生实现,并理解其与排序(sort/argsort)扩展机制的异同。
特性概览:分区操作正式接入 ArrayMethod API
NumPy 的用户自定义 DType(User-Defined DType,以下简称 UDType)机制允许第三方扩展定义全新的数据类型,并借助 ArrayMethod API 为各类 ufunc 与数组方法提供原生实现。此前,自定义 DType 已经可以通过 ArrayMethod API 注册sort与argsort实现(相关细节见 C API 文档中的 Sorting and Argsorting 一节)。
本次新增特性在此基础上进一步扩展:
- 用户自定义 DType 现在可以像实现排序一样,为
partition(就地分区)与argpartition(分区索引)注册自定义实现; - 这些实现通过
numpy.partition与numpy.argpartition(以及对应的ndarray.partition、ndarray.argpartition方法)在传入用户自定义 DType 数组时被自动调用; - 注册入口与排序一致:将实现对应操作的 ArrayMethod spec 传递给
PyUFunc_AddLoopsFromSpecs函数即可,无需额外的专用注册 API。
从源码结构看,该能力由 numpy/_core/src/multiarray/item_selection.c 中的分区分发逻辑支撑,并在 numpy/_core/tests/test_custom_dtypes.py 中有完整的测试覆盖,属于可开箱使用的正式功能,而非实验性接口。
注册入口:PyUFunc_AddLoopsFromSpecs 与 PyUFunc_LoopSlot
注册自定义 DType 的 partition / argpartition 实现,核心入口是PyUFunc_AddLoopsFromSpecs(自 NumPy 2.4 起引入,见 C API 文档)。该函数接收一个以NULL结尾的PyUFunc_LoopSlot数组,每个槽位描述"为哪一个操作注册哪一个 ArrayMethod spec"。
PyUFunc_LoopSlot结构体包含两个成员(定义见 C API 文档):
| 成员 | 类型 | 说明 |
|---|---|---|
name | const char * | 要注册到的操作名称,格式与 entry point 类似:(module ':')? (object '.')* name,默认模块为numpy。例如sin、strings.str_len、numpy.strings:str_len。注意:部分名称并不直接对应 ufunc,如"sort"、"argsort"、"real"、"imag"——它们在内部使用 ufunc 或 ufunc-like 机制实现,"partition"与"argpartition"同样属于这一类。 |
spec | PyArrayMethod_Spec * | 用于创建 loop 的 ArrayMethod spec。 |
分区实现即通过将name设为"partition"或"argpartition"、spec设为对应实现来注册。注册示意如下:
static PyArrayMethod_Spec partition_spec = { .nin = 2, .nout = 1, .dtypes = partition_dtypes, // {你的 DType, NPY_INTP, 你的 DType} .slots = partition_slots, // 例如 NPY_METH_resolve_descriptors 与 NPY_METH_get_loop .flags = NPY_METH_NO_FLOATINGPOINT_ERRORS, }; PyUFunc_LoopSlot loops[] = { {"partition", &partition_spec}, {"argpartition", &argpartition_spec}, {NULL, NULL}, // NULL 结尾 }; int res = PyUFunc_AddLoopsFromSpecs(loops);与 sort / argsort 注册的异同
- 相同点:
"partition"、"argpartition"与"sort"、"argsort"一样,都是通过PyUFunc_AddLoopsFromSpecs批量注册,可与常规 ufunc loops 一起传递; - 不同点:sort 系列的 spec 要求
nin=1, nout=1(单输入单输出,排序就地完成),而 partition 系列的 spec 要求nin=2, nout=1——两个输入分别是被分区数组与分区索引(kth)数组,详见下文契约。
Partition Spec 契约:输入、输出与就地语义
根据 C API 文档 "Partitioning and Argpartitioning" 一节,注册 partition / argpartition 的 ArrayMethod spec 必须遵守以下契约:
nin=2, nout=1:- 第一个输入
data[0]:要分区的数组; - 第二个输入
data[1]:kth 索引数组,即按哪些位置进行分区; - 输出(仅 argpartition):分区后的索引数组。
- 第一个输入
就地分区约束:partition 是就地操作,因此强制要求
data[0] == data[2](第一个输入与输出共享同一块内存/数据指针)。kth 数组的形态约束:
data[1]始终是一个NPY_INTP类型的连续(contiguous)数组,内含分区索引。若传入多个分区索引,则数组会对每个索引依次分区(多个 kth 值会先被排序以确保分区互不干扰)。argpartition 的输出类型:argpartition 返回的是新分配的索引数组,因此输出必须是
NPY_INTP类型。循环维度信息:在 strided loop 中,
dimensions[0]表示要分区的元素个数,dimensions[1]表示分区索引的个数。
与 Python 层语义的对应
上述契约与numpy.partition/numpy.argpartition的既有 Python 语义完全一致:
kth支持负索引(会被转换为shape[axis] + kth);- 索引越界会抛出
ValueError; kind参数接受'introselect'(默认)与'stable'两种选择算法,并可配合descending=True进行降序分区(该参数在 NumPy 2.0 引入的稳定选择机制上提供)。
参数结构:PyArrayMethod_PartitionParameters 与 NPY_SELECTKIND
与 sort 系列通过PyArrayMethod_SortParameters传递排序标志类似,分区操作通过PyArrayMethod_PartitionParameters将分区语义信息传给 loop:
typedef struct { NPY_SELECTKIND flags; } PyArrayMethod_PartitionParameters;- 该结构通过 loop 收到的
context->parameters字段访问(context即PyArrayMethod_Context); flags是NPY_SELECTKIND枚举值的按位或,指示本次分区执行的种类——关键语义是是否为降序分区。
NPY_SELECTKIND定义在 numpy/_core/include/numpy/ndarraytypes.h#L199-L207:
typedef enum { _NPY_SELECT_UNDEFINED = -1, NPY_INTROSELECT = 0, // new style names NPY_SELECT_DEFAULT = 0, NPY_SELECT_STABLE = 2, NPY_SELECT_DESCENDING = 4, } NPY_SELECTKIND; #define NPY_NSELECTS (NPY_SELECT_DESCENDING + 1)即合法取值包括NPY_SELECT_DEFAULT(等价于NPY_INTROSELECT,默认选择算法)、NPY_SELECT_STABLE(稳定分区)与NPY_SELECT_DESCENDING(降序分区),NPY_NSELECTS定义为 5。
编写 loop 的关键建议(来自官方文档):如果 strided loop 的实现依赖flags(例如需要区分升降序),最佳实践是只定义NPY_METH_get_loop槽位,而不设置其他 loop slots。这样可以在get_strided_loop阶段根据context->parameters中的 flags 动态选择/生成合适的内层循环,避免在单一路径中处理所有分支。
源码级实现剖析:分区分发调用链
理解注册契约后,再看底层分发逻辑,可以印证上述约定的每一个细节。核心实现在 numpy/_core/src/multiarray/item_selection.c。
1. kth 数组预处理:partition_prep_kth_array
无论是 partition 还是 argpartition,入口都会先调用partition_prep_kth_array(item_selection.c#L1666-L1720)对 kth 数组做规范化:
- 类型检查:拒绝布尔数组(
ValueError: Booleans unacceptable as partition index),非整数数组抛TypeError: Partition index must be integer; - 维度检查:kth 数组维度必须
<= 1; - 强制转型:将 kth 数组转换为
NPY_INTP类型(PyArray_Cast(ktharray, NPY_INTP)); - 负索引归一化:将负索引加上
shape[axis],并做越界检查(空数组除外),越界抛ValueError: kth(=...) out of bounds (...); - 多索引排序:当 kth 值多于一个时,先对 kth 数组排序,确保多次分区互不干扰——这与文档中"对每个索引依次分区"的契约对应。
2. PyArray_Partition:就地分区分发
PyArray_Partition(item_selection.c#L1727-L1808)的流程清晰展示了 ArrayMethod 的接入方式:
- 校验分区种类
which是否在[0, NPY_NSELECTS)范围内,非法值抛ValueError: not a valid partition kind; - 预处理 kth 数组(见上);
- 从 DType 槽位取方法:
method = NPY_DT_SLOTS(NPY_DTYPE(PyArray_DESCR(op)))->part_meth;——这正是注册"partition"spec 时被填充的槽位; - 回退机制:若
method == NULL(DType 未注册 partition 实现),回退到排序实现:return PyArray_Sort(op, axis, (NPY_SORTKIND)which);,注释明确"Use sorting, slower but equivalent"(用排序代替,较慢但结果等价)。这意味着未注册自定义分区的 DType 依然可用 partition,只是性能退化; - 构造
dtypes[3] = {dt, kdt, dt}与given_descrs[3] = {descr, kdescr, descr}——印证了"data[0] == data[2](输入与输出同类型同描述符)"的契约,其中kdt是预处理后 kth 数组的NPY_INTPDType; - 依次调用
method->resolve_descriptors(...)与method->get_strided_loop(...),将PyArrayMethod_PartitionParameters通过context.parameters传入; - 由于分区对每个轴切片都是连续的,strides 直接按
loop_descrs[i]->elsize构造(注释:"Arrays are always contiguous for partitioning"); - 最终交给
_new_sortlike执行多轴循环。
3. PyArray_ArgPartition:索引分区分发
PyArray_ArgPartition(item_selection.c#L1815-L1909)与 partition 几乎对称,差异在于:
- 方法取自
NPY_DT_SLOTS(...)->argpart_meth槽位; - 输出描述符固定为
odescr = PyArray_DescrFromType(NPY_INTP);,dtypes 变为{dt, kdt, odt}——印证了"输出必须是NPY_INTP"的契约; - 回退逻辑改为调用
PyArray_ArgSort(op2, axis, (NPY_SORTKIND)which); - 最终交给
_new_argsortlike执行。
4. 槽位与分发的对应关系
从 numpy/_core/code_generators/numpy_api.py 与 numpy/_core/src/umath/dispatching.cpp 中PyUFunc_AddLoopsFromSpecs的实现来看,注册系统在解析"partition"/"argpartition"名称时会将其路由到 DType 的part_meth/argpart_meth槽位,与"sort"/"argsort"的sort_meth/argsort_meth机制同构。这也解释了为什么文档明确提示"以与排序和 argsorting 类似的方式"实现。
降序分区的处理:以 test_custom_dtypes 为参考
仓库中的自定义 DType 测试类(numpy/_core/tests/test_custom_dtypes.py#L448-L512)对 partition 与 argpartition 的注册实现做了全面验证,是学习编写实现的极佳范本。其测试要点包括:
test_partition:构造以不同缩放因子表示的 float64 视图 DType 数组(正序、逆序、非对齐 stride 等形态),调用a.partition(k)后,断言a[:k]与a[k:]各自排序后的内容符合预期;同时覆盖descending=True的降序分区(断言分区点前后元素集合互换);test_argpartition:对相同数组形态调用a.argpartition(k),通过返回索引indices回取元素并验证分区正确性,同样覆盖降序模式。
这两个测试同时验证了:
- 自定义 partition 实现确实被
ndarray.partition/ndarray.argpartition分发调用; - 就地语义(partition 直接修改原数组)与
NPY_INTP输出类型约束成立; - 非连续(stride 反向)与非对齐数组下实现依然正确。
版本与适用性说明
PyUFunc_AddLoopsFromSpecs自NumPy 2.4起引入(.. versionadded:: 2.4),本特性为后续迭代中为"partition"/"argpartition"扩展的支持;- 官方 C API 文档(doc/source/reference/c-api/array.rst)是了解 ArrayMethod spec 各字段(
resolve_descriptors、get_strided_loop、PyArrayMethod_Spec、PyArrayMethod_Context等)的一手资料; - 当前仓库的 2.5.0 发布说明(doc/source/release/2.5.0-notes.rst)还显示,类型标注层面已为
argpartition提供 shape 相关的返回类型推断(shape-typing),说明该 API 已被视为稳定的公开能力并持续获得工具链完善。
小结
DType 的 partition / argpartition ArrayMethod 支持,为自定义数据类型在保持 NumPy 统一 API 语义的前提下接入高性能原生分区算法铺平了道路。核心要点可归纳为四条:
- 通过
PyUFunc_AddLoopsFromSpecs以"partition"/"argpartition"为名注册 spec,与 sort / argsort 共用同一注册机制; - 严格遵守
nin=2, nout=1、data[0] == data[2](就地)、data[1]为NPY_INTP连续数组、argpartition 输出为NPY_INTP的契约; - 借助
context->parameters中的PyArrayMethod_PartitionParameters.flags(NPY_SELECTKIND)判断升降序,推荐只定义NPY_METH_get_loop槽位; - 未注册实现的 DType 会自动回退到排序实现(结果等价、性能较慢),因此该特性是完全向后兼容的可选优化。
【免费下载链接】numpyThe fundamental package for scientific computing with Python.项目地址: https://gitcode.com/gh_mirrors/nu/numpy
创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考