
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_knp.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.pydef top_k(a, k, /, *, axis-1, modelargest, sortedTrue):与之对应的类型桩type stub在 numpy/_core/fromnumeric.pyidef top_k( a: ArrayLike, k: int, /, *, axis: int -1, mode: Literal[largest, smallest] largest, sorted: bool True, ) - tuple[NDArray[Any], NDArray[intp]]: ...各参数语义如下表参数类型默认值说明aarray_like必填源数组支持任意可被asanyarray转换的输入含列表、嵌套列表、字符串数组等见后文测试部分kint必填需要返回的最大/最小元素个数。必须是非负整数且不能超出axis指定的维度大小axisint-1沿哪个轴查找最大/最小元素默认是最后一个轴。不支持None与sort/argsort不同top_k不允许传入None表示展平操作mode{largest,smallest}largestlargest返回最大的 k 个元素smallest返回最小的 k 个元素sortedboolTrue为True时返回的 k 个元素保证按降序largest或升序smallest排序为False时不保证有序但计算上更宽松需要注意两个“仅关键字”设计位置参数之后紧跟/表示a与k只能按位置传入axis、mode、sorted前有*表示这三个参数必须用关键字方式传入如np.top_k(a, 2, axis0)不能写成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)的二维数组沿axis1取 k 个元素返回的两个数组形状都是(m, k)沿axis0取则形状都是(k, n)。返回值可直接用于“取值并验证”的闭环操作用np.take_along_axis(a, topk_indices, axisaxis)可以还原出topk_values这正是测试代码中使用的等价性断言见 numpy/_core/tests/test_multiarray.pyx_value, x_ind np.top_k(a, k, axisaxis, modemode, sortedsorted) assert_equal(np.take_along_axis(a, x_ind, axisaxis), x_value)四、使用示例以下示例完整取自函数 docstringnumpy/_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, axis0) (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 个modesmallest np.top_k(a, 2, axis1, modesmallest) (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。换言之modelargest时 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], dtypedtype) val, ind np.top_k(a, 3) assert not np.isnan(val).any()5.2 索引稳定性不保证稳定文档注释强调返回的索引不保证稳定即对于重复值返回索引的顺序与它们在输入数组中的出现顺序不一定一致——这一约束与sorted参数取值无关。因此在需要精确的“第一个出现位置”语义时不应依赖top_k的索引顺序。5.3sortedFalse只影响输出顺序不影响取值集合sortedFalse表示结果不保证有序但返回的仍然是“某 k 个最大/最小元素”这一集合。测试通过“先取后排序再比较”的方式对两种模式分别校验numpy/_core/tests/test_multiarray.py保证无论sorted取何值最终得到的值集合与索引集合与参考结果一致。5.4 非法参数的错误提示实现中对三类非法输入显式抛出ValueError测试逐一覆盖numpy/_core/tests/test_multiarray.py非法输入错误信息k 0如k-2k(-2) provided must be a non-negative integer.modeinvalid等非法取值mode(invalid) must be either largest or smallest.axisNoneaxisNone is not supported. Please provide a valid axis.六、源码实现原理argpartitiontake_along_axistop_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([], dtypenp.intp) indices np.argpartition(arr, kth, axisaxis, descendinglargest) slice_ (np.s_[:],) * axis (np.s_[:k],) indices indices[slice_] values np.take_along_axis(arr, indices, axisaxis) if sorted: sort_indices np.argsort(values, axisaxis, descendinglargest, stableFalse) values np.take_along_axis(values, sort_indices, axisaxis) indices np.take_along_axis(indices, sort_indices, axisaxis)各步骤的工程含义如下输入归一化np.asanyarray(a)将任意array_like统一为 ndarray或子类normalize_axis_index负责把负轴如-1规约为非负索引并校验越界。部分选择而非全排序使用np.argpartition(arr, kth, axisaxis, descendinglargest)仅做“划分”式选择——这正是top_k的性能来源。它并不对整条轴排序而是把第k-1大的元素放到划分点保证左侧largest模式即为前 k 个候选。k0时构造空索引数组返回空切片。截取前 k通过构造多维切片slice_axis之前各维全取:目标轴取:k只保留划分后前 k 个位置。取值与可选的排序用np.take_along_axis依据索引取出对应值若sortedTrue再对值做一次np.argsortdescendinglargest、stableFalse并把排序顺序同时应用到values与indices从而保证“值有序索引随之对齐”。从实现可以推断top_k的时间复杂度主要由argpartition主导平均 O(n) 级别的选择开销而非全排序的 O(n log n)仅当sortedTrue时才额外对 k 个元素做小规模排序。对于k远小于轴长度的“大数组取前几”场景这是比“先np.sort再切片”更省的做法。另外值得注意top_k借助array_function_dispatch机制注册了派发器_top_k_dispatchernumpy/_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 引入的字符串 dtypedtypeT。测试 numpy/_core/tests/test_stringdtype.py 验证了字符串数组上的largest与smallest两种模式def test_top_k(string_list): arr np.array(string_list, dtypeT) expected sorted(string_list, reverseTrue)[: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, modesmallest) 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}覆盖k0、axis-1/1/0、modesmallest以及三类非法参数负 k、非法 mode、axisNone的错误信息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且不需要全局有序时sortedFalse配合划分式选择能避免全量排序开销与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),仅供参考