NumPy 新增numpy.top_k函数:沿指定轴高效提取最大/最小 k 个元素及其索引
【免费下载链接】numpyThe fundamental package for scientific computing with Python.项目地址: https://gitcode.com/gh_mirrors/nu/numpy
导读
本文基于 NumPy 仓库的发布变更记录 doc/release/upcoming_changes/31659.new_function.rst 展开,介绍新增的公共 APInumpy.top_k:它能够在指定轴上一次性返回数组的最大/最小 k 个元素及其对应的索引,适用于 Top-K 筛选、排序统计、推荐系统候选召回等场景。读完本文,你将掌握top_k的完整签名、参数语义、返回值结构、NaN 与重复值等边界行为,并能结合源码理解其基于argpartition的底层实现原理与既有测试覆盖。
一、变更记录与新增 API 概览
在 NumPy 的发布变更目录doc/release/upcoming_changes/中,编号 31659 的条目记录了本次新增的函数:
新增函数
numpy.top_k:np.top_k(array, k, axis=..., mode=..., sorted=...),沿给定轴返回数组中最大/最小的 k 个值。
与以往需要组合np.argpartition、np.argsort、np.take_along_axis才能完成的“取 Top-K 及其索引”操作不同,top_k将这一流程封装为单一、稳定的公共接口,且直接导出到numpy顶层命名空间(参见 numpy/init.py 的导入以及 numpy/_core/init.pyi 与 numpy/_core/numeric.pyi 中的导出声明)。
该函数定位于fromnumeric模块(与sort、take、argpartition等归置同一文件),是标准的“array_function_dispatch”风格公共函数。
二、函数签名与参数详解
完整的函数签名定义在 numpy/_core/fromnumeric.py:
def top_k(a, k, /, *, axis=-1, mode="largest", sorted=True):与之对应的类型桩(type stub)在 numpy/_core/fromnumeric.pyi:
def top_k( a: ArrayLike, k: int, /, *, axis: int = -1, mode: Literal["largest", "smallest"] = "largest", sorted: bool = True, ) -> tuple[NDArray[Any], NDArray[intp]]: ...各参数语义如下表:
| 参数 | 类型 | 默认值 | 说明 |
|---|---|---|---|
a | array_like | 必填 | 源数组,支持任意可被asanyarray转换的输入(含列表、嵌套列表、字符串数组等,见后文测试部分) |
k | int | 必填 | 需要返回的最大/最小元素个数。必须是非负整数,且不能超出axis指定的维度大小 |
axis | int | -1 | 沿哪个轴查找最大/最小元素,默认是最后一个轴。不支持None(与sort/argsort不同,top_k不允许传入None表示展平操作) |
mode | {"largest","smallest"} | "largest" | "largest"返回最大的 k 个元素,"smallest"返回最小的 k 个元素 |
sorted | bool | True | 为True时,返回的 k 个元素保证按降序(largest)或升序(smallest)排序;为False时不保证有序,但计算上更宽松 |
需要注意两个“仅关键字”设计:
- 位置参数之后紧跟
/,表示a与k只能按位置传入; axis、mode、sorted前有*,表示这三个参数必须用关键字方式传入,如np.top_k(a, 2, axis=0),不能写成np.top_k(a, 2, 0)。
类型桩中k: int且返回值标注为tuple[NDArray[Any], NDArray[intp]],其中NDArray[intp]说明索引数组为平台相关的整数指针类型(64 位平台上即int64)。
三、返回值:(values, indices)二元组
top_k返回的是包含两个数组的元组:
topk_values:前 k 个值组成的数组;topk_indices:与topk_values一一对应的索引数组。
两个数组的形状均为:输入数组的形状,将axis那一维替换为k。也就是说,对于形状为(m, n)的二维数组沿axis=1取 k 个元素,返回的两个数组形状都是(m, k);沿axis=0取,则形状都是(k, n)。
返回值可直接用于“取值并验证”的闭环操作:用np.take_along_axis(a, topk_indices, axis=axis)可以还原出topk_values,这正是测试代码中使用的等价性断言(见 numpy/_core/tests/test_multiarray.py):
x_value, x_ind = np.top_k(a, k, axis=axis, mode=mode, sorted=sorted) assert_equal(np.take_along_axis(a, x_ind, axis=axis), x_value)四、使用示例
以下示例完整取自函数 docstring(numpy/_core/fromnumeric.py),可直接在交互环境中复现。
4.1 默认行为:沿最后一个轴取最大的 k 个
>>> import numpy as np >>> a = np.array([[1, 2, 3, 4, 5], [5, 4, 3, 2, 1]]) >>> np.top_k(a, 2) (array([[5, 4], [5, 4]]), array([[4, 3], [0, 1]]))第一行[1, 2, 3, 4, 5]中最大的两个是5, 4,索引为4, 3;第二行[5, 4, 3, 2, 1]中最大的两个同样是5, 4,索引为0, 1。
4.2 沿第 0 轴取最大的 k 个
>>> np.top_k(a, 2, axis=0) (array([[5, 4, 3, 4, 5], [1, 2, 3, 2, 1]]), array([[1, 1, 0, 0, 0], [0, 0, 1, 1, 1]]))此时按列比较两行:例如第 0 列两个元素1与5中,最大值5来自第 1 行(索引1),次大值1来自第 0 行(索引0),输出第 0 列结果[5, 1]及其索引[1, 0]。
4.3 取最小的 k 个(mode="smallest")
>>> np.top_k(a, 2, axis=1, mode="smallest") (array([[1, 2], [1, 2]]), array([[0, 1], [4, 3]]))4.4 含 NaN 的浮点数组
>>> np.top_k(np.array([1., 2., 3., np.nan]), 2) (array([3., 2.]), array([2, 1]))NaN 被排到末尾,因此当 k 小于数组中 NaN 出现位置之后的有效元素数量时,返回结果中不会出现 NaN(详见下一节)。
五、关键行为与边界语义
5.1 NaN 处理:与sort一致,NaN 排在末尾
文档明确指出:与排序行为类似,NaN 会被推到末尾,因此只有当 NaN 恰好落在前 k 个位置(即数组中 NaN 过多)时,它们才会出现在输出中——无论mode是"largest"还是"smallest"。换言之,mode="largest"时 NaN 不会挤占正常数值的前 k 位置。
该行为由测试 numpy/_core/tests/test_multiarray.py 专门验证,且覆盖了 NumPy 全部浮点类型码(np.typecodes["AllFloat"],包含半精度、单精度、双精度、扩展精度与复数类型):
@pytest.mark.parametrize("dtype", np.typecodes["AllFloat"]) def test_top_k_floating_nan(self, dtype): a = np.array([np.nan, 1, 2, 3, np.nan], dtype=dtype) val, ind = np.top_k(a, 3) assert not np.isnan(val).any()5.2 索引稳定性:不保证稳定
文档注释强调:返回的索引不保证稳定,即对于重复值,返回索引的顺序与它们在输入数组中的出现顺序不一定一致——这一约束与sorted参数取值无关。因此在需要精确的“第一个出现位置”语义时,不应依赖top_k的索引顺序。
5.3sorted=False只影响输出顺序,不影响取值集合
sorted=False表示结果不保证有序,但返回的仍然是“某 k 个最大/最小元素”这一集合。测试通过“先取后排序再比较”的方式对两种模式分别校验(numpy/_core/tests/test_multiarray.py),保证无论sorted取何值,最终得到的值集合与索引集合与参考结果一致。
5.4 非法参数的错误提示
实现中对三类非法输入显式抛出ValueError,测试逐一覆盖(numpy/_core/tests/test_multiarray.py):
| 非法输入 | 错误信息 |
|---|---|
k < 0(如k=-2) | k(=-2) provided must be a non-negative integer. |
mode="invalid"等非法取值 | mode(="invalid") must be either "largest" or "smallest". |
axis=None | axis=None is not supported. Please provide a valid axis. |
六、源码实现原理:argpartition+take_along_axis
top_k的核心实现位于 numpy/_core/fromnumeric.py,整体是一个清晰的四步流水线:
arr = np.asanyarray(a) axis = normalize_axis_index(axis, arr.ndim) kth = k - 1 if k > 0 else np.array([], dtype=np.intp) indices = np.argpartition(arr, kth, axis=axis, descending=largest) slice_ = (np.s_[:],) * axis + (np.s_[:k],) indices = indices[slice_] values = np.take_along_axis(arr, indices, axis=axis) if sorted: sort_indices = np.argsort(values, axis=axis, descending=largest, stable=False) values = np.take_along_axis(values, sort_indices, axis=axis) indices = np.take_along_axis(indices, sort_indices, axis=axis)各步骤的工程含义如下:
- 输入归一化:
np.asanyarray(a)将任意array_like统一为 ndarray(或子类);normalize_axis_index负责把负轴(如-1)规约为非负索引,并校验越界。 - 部分选择而非全排序:使用
np.argpartition(arr, kth, axis=axis, descending=largest)仅做“划分”式选择——这正是top_k的性能来源。它并不对整条轴排序,而是把第k-1大的元素放到划分点,保证左侧(largest模式)即为前 k 个候选。k=0时构造空索引数组,返回空切片。 - 截取前 k:通过构造多维切片
slice_(axis之前各维全取:,目标轴取:k),只保留划分后前 k 个位置。 - 取值与可选的排序:用
np.take_along_axis依据索引取出对应值;若sorted=True,再对值做一次np.argsort(descending=largest、stable=False),并把排序顺序同时应用到values与indices,从而保证“值有序,索引随之对齐”。
从实现可以推断:top_k的时间复杂度主要由argpartition主导(平均 O(n) 级别的选择开销,而非全排序的 O(n log n)),仅当sorted=True时才额外对 k 个元素做小规模排序。对于k远小于轴长度的“大数组取前几”场景,这是比“先np.sort再切片”更省的做法。
另外值得注意:top_k借助array_function_dispatch机制注册了派发器_top_k_dispatcher(numpy/_core/fromnumeric.py),因此对实现__array_function__协议的第三方数组库,np.top_k也可被正确分派。
七、类型标注与更广泛的输入支持
7.1 类型桩
函数在 numpy/_core/fromnumeric.pyi 中提供了完整类型声明,返回值被精确标注为tuple[NDArray[Any], NDArray[intp]],且mode使用Literal["largest", "smallest"]枚举约束,sorted默认为True。同时,top_k已加入numpy/_core/__init__.pyi与numpy/_core/numeric.pyi的导出列表(__all__),确保使用类型检查工具(如 mypy、pyright)时能获得完整的补全与校验。
7.2 支持新字符串 dtype
除数值数组外,top_k同样适用于 NumPy 2.x 引入的字符串 dtype(dtype="T")。测试 numpy/_core/tests/test_stringdtype.py 验证了字符串数组上的largest与smallest两种模式:
def test_top_k(string_list): arr = np.array(string_list, dtype="T") expected = sorted(string_list, reverse=True)[:2] values, indices = np.top_k(arr, 2) assert values.tolist() == expected assert arr[indices].tolist() == expected expected = sorted(string_list)[:2] values, indices = np.top_k(arr, 2, mode="smallest") assert values.tolist() == expected assert arr[indices].tolist() == expected7.3 接受 Python 列表等非数组输入
在 numpy/_core/tests/test_numeric.py 中,top_k直接以嵌套列表作为输入并返回期望结果,印证了asanyarray归一化对array_like输入的通用支持。
八、测试覆盖总结
top_k在仓库中拥有成体系的测试矩阵,可作为使用时的行为契约参考:
| 测试位置 | 覆盖点 |
|---|---|
| numpy/_core/tests/test_multiarray.py | 参数化sorted ∈ {True, False};覆盖k=0、axis=-1/1/0、mode="smallest",以及三类非法参数(负 k、非法 mode、axis=None)的错误信息 |
| numpy/_core/tests/test_multiarray.py | 全部浮点类型下 NaN 被推至末尾、不进入前 k 结果 |
| numpy/_core/tests/test_numeric.py | 嵌套列表输入的基本正确性 |
| numpy/_core/tests/test_stringdtype.py | 新字符串 dtype 上两种模式的取值与索引还原 |
其中assert_top_k辅助方法(numpy/_core/tests/test_multiarray.py)通过“索引取值还原 + 排序后与参考结果比对”的方式,从两个独立维度交叉验证了返回值的一致性,这也为使用者提供了自测同类逻辑的参考范式。
九、典型应用场景
- Top-K 召回与筛选:在推荐、检索场景中沿批量维度(如
axis=-1)一次性取出每条样本得分最高的 k 个候选及其位置,避免手写argpartition+ 切片 + 取值三板斧; - 统计分析:需要同时获得极值集合与其位置(如找出 k 个最大异常点及其下标)时,
(values, indices)二元组可直接消费; - 内存/时间敏感的排序替代:当
k << n且不需要全局有序时,sorted=False配合划分式选择能避免全量排序开销; - 与
argpartition/sort互补:argpartition只返回索引且不保证有序,sort做全量排序,top_k位于两者之间——一次调用同时给出有序值与索引,属于高层封装接口。
结语
numpy.top_k是 NumPy 在既有排序/划分原语之上新增的面向任务型编程的公共函数,用一处调用替代了以往多步组合操作,并完整覆盖了轴方向、最大/最小模式、有序/无序输出、NaN 语义与非法参数校验等细节。其实现建立在成熟的argpartition与take_along_axis机制之上,类型桩、导出声明与多组测试均已齐备,可作为日常数据筛选与 Top-K 分析的首选入口。如需深入研读,可从入口实现 numpy/_core/fromnumeric.py 及其配套测试开始。
【免费下载链接】numpyThe fundamental package for scientific computing with Python.项目地址: https://gitcode.com/gh_mirrors/nu/numpy
创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考