1. 为什么你需要“最小的k个元素”
1.1 三个最典型的应用场景
先说个实际背景。我在做推荐系统的时候,经常要对一批候选物品的得分做筛选,得分最低的那几个往往代表用户最不感兴趣的内容,需要拿来当负样本。这类操作在算法里出现频率极高,但很多同学第一反应是 torch.sort 排个序,再切片取前 k 个。功能上没问题,但性能和代码优雅度都差点意思。
用到“获取最小的k个元素”的场景,我随手就能列三个:
- KNN 最近邻:给定查询向量和一堆候选向量,求距离最小的 k 个邻居。这里距离就是核心排序依据,取前 k 小是刚需。
- 难样本挖掘:在度量学习或对比学习里,经常要挑最不相似的同类样本(或者其他难例),本质上也是取某个相似度/距离度量最小的 k 个。
- 强化学习 / 采样处理:某些策略里要从动作价值估计中排除最差的动作,或者在 beam search 里保留最小惩罚路径,这些都要拿最小的 k 个索引。
你会发现,“取最小 k 个”不是一个孤立操作,它往往和后续的索引重排、mask 生成、gather 取值绑在一起。所以你需要的不是一个“能用就行”的写法,而是一个稳定、高效、能直接返回索引的工具。
1.2 方案选型:为什么首选 torch.topk
先列一下常见方案,再逐个说它们的问题。
| 方案 | 核心思路 | 时间复杂度 | 返回索引 | 缺点 |
|---|---|---|---|---|
torch.sort + 切片 |
全排序后取前 k | O(n log n) | 需额外处理 | 多算了很多用不上的排序 |
循环 torch.min |
每次取最小,mask 掉后再取 | O(kn) | 需自己维护索引 | 慢,代码笨重 |
torch.topk(largest=False) |
部分排序取 TopK | O(n log k) | 直接返回 values 和 indices | 需要理解参数语义 |
sorted=True 时,topk 内部实际上做了一个类似部分排序的优化,返回的结果是排好序的,但你只花了 O(n log k) 的成本,而不是把整组数据全部排完。数据量越大,这个差距越明显。sort 简单直观,但代价是你把已经排好的后面 n-k 个元素也算了一遍,纯属浪费。
torch.min 循环更不用提——每轮都要维护一个 mask,把上一轮最小值遮掉,再跑一次全量扫描。如果 k 是几十几百,张量又大,这个循环的耗时非常可观,放在训练步骤里基本没法忍。
所以在 PyTorch 里,“获取最小的k个元素”这件事的标准答案就是 torch.topk(input, k, dim=..., largest=False)。它快、支持 GPU、返回索引,还能指定维度,几乎是为这个需求量身定制的。
1.3 topk 为什么不走全排序
很多人好奇,为什么 topk 比 sort 快那么多。核心原因是它不需要把整个序列排好序。你可以把 topk 内部理解成一个“部分排序”过程:它只关心边界那一块,也就是前 k 个(或后 k 个)的集合,集合内部的顺序只在 sorted=True 时才额外维护。
类比一下生活场景:你从 10000 份简历里挑出最差的 10 份扔掉,根本不需要把这 10000 份按好坏排成一条完整的长队。你只要一直维护一个“当前最差的 10 份”的集合,遇到更差的就替换掉集合里最“好”的那个。这样每份简历只需要和集合里的 10 个比较,而不是每次都要排全队。这就是 O(n log k) 和 O(n log n) 的直观区别。
PyTorch 底层在 GPU 上对这个操作有专门优化,实测下来在大规模张量上,topk 和 sort 的耗时可差好几倍。这一点在后面的性能对比部分我会放具体数字。
需要模型API调用? 免费领10W Token,多模型网关一键接入 Claude、DeepSeek 等主流模型。
2. 核心 API 解析与参数避坑指南
2.1 torch.topk 的完整签名和参数语义
先看官方签名:
python复制torch.topk(input, k, dim=None, largest=True, sorted=True, *, out=None)
每个参数的含义:
input:输入张量,任意维度。k:要取的元素个数,必须大于等于 1,且不能超过指定维度上的长度。dim:沿着哪个维度取。默认是-1,也就是最后一个维度。largest:True表示取最大的 k 个,False表示取最小的 k 个。这里就是“获取最小的k个元素”的关键开关。sorted:返回结果是否排序。True时结果按从小到大(largest=False)或从大到小(largest=True)排好。False时不保证顺序,但速度可能更快,内存开销也稍低。out:可选的输出元组,一般用不到。
我第一次用的时候,只记住了 torch.topk(x, 3),结果返回的是最大的 3 个,搞了半天才发现忘记设 largest=False。这个参数太容易漏了,建议写成关键字参数 largest=False,不要省。
2.2 返回值的正确打开方式
torch.topk 返回一个 (values, indices) 的命名元组:
values:选出来的 k 个元素的值,形状是原张量在指定维度上变成 k。indices:这些元素在原张量中的索引,形状和values一模一样,类型是torch.int64。
拿一维张量举例:
python复制import torch
x = torch.tensor([5.0, 2.0, 8.0, 1.0, 3.0])
values, indices = torch.topk(x, 2, largest=False)
# values: tensor([1.0, 2.0])
# indices: tensor([3, 1])
这里 indices 告诉你最小的两个元素在原张量里的位置是索引 3 和 1。有了索引,你就可以用 torch.gather 或者直接按索引去原张量里做后续操作,比如给这些位置赋值、生成 mask 等。
多维张量时,indices 里存的是对应维度上的索引,不是全局索引。这个非常容易搞混,我后面第 4 章会专门讲这个坑。
2.3 边界条件与异常行为
几个边界情况你必须提前知道,否则会上线才发现问题:
- k 不能为 0:
torch.topk(x, k=0)会直接抛运行时错误。虽然语义上取 0 个好像没问题,但 PyTorch 没有支持这个边界,代码里要注意判断。 - k 不能超过维度长度:在长度为 5 的维度上取 k=6,会报错
k larger than dimension。 - NaN 的影响:张量里如果有
NaN,排序结果会非常诡异,NaN可能出现在任何位置,而且不保证稳定。取最小 k 个时,NaN很可能混进来,导致你拿到一堆没意义的值。 - 空张量:如果输入张量某个维度长度为 0,取 topk 会报错,需要先判断空数据。
这些边界问题看起来不起眼,但在数据清洗、异常输入时非常常见。我的习惯是封装一层工具函数,统一处理这些边界,而不是在业务代码里到处判断。
3. 从一维到多维:实操代码与性能表现
3.1 一维张量取前 k 小:最基础用法
一维场景是最直观的,直接上代码:
python复制import torch
scores = torch.randn(10000)
k = 5
values, indices = torch.topk(scores, k, largest=False)
print(values)
print(indices)
# 验证:取出来的值确实是最小的 k 个
sorted_scores, _ = torch.sort(scores)
assert torch.equal(values, sorted_scores[:k])
这里我加了一个断言,用 sort 的结果做验证。实际项目里这种验证会帮你在初期排查掉很多问题,特别是你刚把一个算法写出来、不确定 dim 和 largest 配没配对的时候。
要注意的是,topk 返回的 values 顺序是升序(因为 largest=False 且默认 sorted=True)。如果你不关心顺序,可以设置 sorted=False,理论上有一点点性能收益,实测不太明显,但如果有同学在追求极致的推理速度,可以打开。
3.2 多维张量沿 dim 取前 k 小
二维场景就开始有迷惑性了。假设你有一个矩阵 x,形状是 (batch_size, num_features),你想对每个样本分别取最小的 3 个特征值:
python复制x = torch.randn(4, 10) # 4 个样本,每个 10 个特征
values, indices = torch.topk(x, k=3, dim=1, largest=False)
# values 形状: (4, 3)
# indices 形状: (4, 3)
此时 indices[i] 里的数字表示第 i 个样本中,这 3 个最小值所在的特征位置。这里容易出错的地方是 dim=1 写错,很多人会默认对最后一维操作而漏掉 dim。如果张量是 (seq_len, batch, hidden) 这类布局,漏掉 dim 就会默认在最后一维取,结果完全不是你想要的。
三维以上的场景同样适用,dim 指定哪个轴,就从哪个轴取。需要注意的是,indices 只记录指定维度上的位置,其他维度保持不变。也就是说,values 和 indices 的形状,是在原形状基础上把 dim 那维换成 k。
3.3 批量任务如何高效取最小 k 个
很多任务不是一次性取完,而是要在 batch 里各取各的。例如在对比学习里,损失函数计算时需要对每个样本找它跟所有负样本里相似度最低的 k 个负样本,作为“最难负样本”参与训练。这种场景 topk 天然支持 batch 维度:
python复制# 假设 similarity 形状 (batch, num_negatives)
# 对每个样本取相似度最小的 k 个负样本索引
sim = torch.randn(32, 100)
min_values, min_indices = torch.topk(sim, k=10, dim=1, largest=False)
对应地,如果你要利用这些索引去另一个张量里取值,就要配合 torch.gather 或者高级索引。这里有个小技巧:先构造 batch 的基准索引,再横向拼接目标维度索引,组成完整索引对。
python复制batch_idx = torch.arange(32).unsqueeze(1).expand(32, 10)
selected = original_feature[batch_idx, min_indices]
这种方式比循环快得多,而且代码简洁。
3.4 性能对比:topk 与 sort 的差距
我在一台普通 GPU 上做过简单测试,输入是形状 (10000, 100) 的随机张量,沿最后一维取最小 10 个元素。结果大致如下:
| 方法 | 耗时(相对值) |
|---|---|
torch.sort + 切片 |
约 1.0x |
torch.topk(largest=False) |
约 0.4x |
循环 torch.min |
约 12x 以上 |
也就是说,topk 比 sort 快一倍多,比循环快一个数量级以上。随着张量规模增大,sort 和 topk 的差距会进一步拉大。如果放在训练循环里,每一步都省下这么一笔,整体收益非常明显。
另外提醒一句:topk 在 CPU 和 GPU 上的性能表现不一样,如果张量不大,GPU 反而可能因为 kernel 启动开销更慢。小张量直接用 sort 问题不大,别为了优化而优化。
4. 常见问题与排查技巧实录
4.1 dim 分不清导致结果完全错误
这是我在项目中遇到频率最高的错误。有一个例子让我印象很深:一个序列模型输出形状是 (batch, seq_len, num_classes),我当时想对每个时间步取最小概率的类别,脑子里想着“取最后一维”,就直接写了 torch.topk(probs, 1),完全没问题。但后来需求改成“对每个样本取所有时间步里最小的那个概率”,而我还在沿用默认参数,结果取出来的是每个时间步最小的类别概率,而不是整个序列最小的概率,整个实验结论都变了。
排查这种问题最快的办法是打印形状:
python复制values, indices = torch.topk(x, k, dim=1, largest=False)
print(values.shape, indices.shape)
如果你发现 values.shape 和自己预期不符,第一时间检查 dim。
4.2 indices 的索引对齐问题
indices 返回的是指定维度上的索引,这个索引不能直接拿来全局索引原张量。举个例子:三维张量 (batch, seq, hidden),沿 dim=1 取最小 k 个,得到的 indices 只能告诉你在 seq 这个维度上的位置,实际取值需要把 batch 和 hidden 维度也考虑进去。
我用 torch.gather 做过一次很隐蔽的错误:想从原始张量中取回最小值,写了 torch.gather(x, 1, indices.unsqueeze(-1).expand(...)),结果因为维度没对齐,程序没报错但结果完全错位。最后靠人工抽查几个样本才发现。强烈建议在操作后加一步“还原验证”,把取出来的值重新放回原张量,和原始数据对比一下。
4.3 NaN 和非法输入导致的结果异常
前面提过,张量里一旦有 NaN,topk 的行为就不稳。因为 NaN 和任何数比较都是 False,排序算法对它的处理取决于底层实现,结果可能与平台相关。
我遇到过一次线上推理结果异常,排查到最终发现是输入里混了 NaN,导致取出来的“最小值”里面有 NaN,进一步污染了后续的归一化计算。解决办法是提前做输入清洗:
python复制if not torch.isfinite(x).all():
x = torch.nan_to_num(x, nan=-1e9)
在取最小 k 个时,如果你希望忽略 NaN,一个常见做法是把 NaN 替换成一个极大值,这样它就永远不会被选进最小值集合;如果你想取最大值,就替换成极小值。这个思路简单可靠,比逐元素 mask 高效得多。
4.4 k 值超界与空张量处理
当 k 大于维度长度,topk 会直接报错。业务代码里如果 k 是根据数据动态算出来的,就必须加保护:
python复制def safe_topk_min(x, k, dim=-1):
if x.numel() == 0 or k <= 0:
return x, torch.empty(0, dtype=torch.long, device=x.device)
dim_size = x.size(dim)
k = min(k, dim_size)
return torch.topk(x, k, dim=dim, largest=False)
这种工具函数写好后,后续所有调用处都走它,能避免很多“小线上事故”。我自己的代码库里基本都会有这么一层封装,省心。
4.5 梯度回传时的注意事项
values 是支持梯度的,可以直接参与损失计算;indices 是整数类型,本身不可导。如果你的损失函数里需要对“最小值”这几个元素的位置做梯度更新,那 indices 只能用来取值,不能拿它当可导参数。
有一个常见的操作是“soft TopK”,用某种平滑方式近似 k 个最小值的选取,从而让梯度能更充分流动。目前 PyTorch 官方没有直接的 soft topk API,但你可以用带温度的 softmax 对权重做近似,或者用 straight-through estimator 的思路,把 topk 的索引当作离散决策,梯度从值部分回流。
第 4.5 节这种场景一般出现在算法研究里,普通业务代码很少碰到。但如果你在做模型训练时发现某一步梯度是 0,先检查一下是否把操作落在了 indices 上。
5. 扩展玩法:把 topk 用到实战项目里
5.1 用最小 k 个元素做难样本挖掘
以度量学习中的三元组损失为例,通常需要挑选“最难正样本对”和“最难负样本对”。最难的负样本是距离最近(相似度最高)的负样本,但如果你要做“最不像的同类”筛选,就可以用最小相似度挑出难正样本。
python复制# anchor 和 positive 的相似度矩阵,shape (N, N)
sim_matrix = torch.matmul(anchor, positive.T)
# 每个 anchor 取相似度最小的 5 个 positive 作为难例
_, hard_indices = torch.topk(sim_matrix, k=5, dim=1, largest=False)
这种写法简单高效,比“先全排序再切”清晰得多,而且索引直接拿来做后续张量索引。
5.2 注意力掩码中的 topk 技巧
在注意力机制里,有时候要限制每个 query 只关注最相关的 k 个 key。这时如果关注度定义是“越小越好”(比如某种距离度量),你就要取最小的 k 个注意力权重,然后把其他位置mask掉。
python复制attn_weights = torch.randn(batch, num_heads, q_len, k_len)
_, topk_indices = torch.topk(attn_weights, k=8, dim=-1, largest=False)
mask = torch.zeros_like(attn_weights, dtype=torch.bool)
mask.scatter_(-1, topk_indices, True)
# 只保留最小的 8 个权重,其余置为 -inf
masked_weights = attn_weights.masked_fill(~mask, float('-inf'))
这里 scatter_ 配 topk 的索引,一句就能把 mask 搭好,非常顺。注意 topk_indices 必须和 mask 的维度匹配,实操时先 expand 再 scatter 也是常见写法。
5.3 反向取数:argsort 与 gather 的配合
如果你需要的不是“最小的 k 个值”,而是“剩余的那些”,可以先用 topk 拿到要保留的索引,再用 torch.arange 生成全集索引做差集。这个技巧在删除异常点时很有用。
python复制all_idx = torch.arange(scores.size(-1), device=scores.device)
# 找出最小的 k 个索引
_, min_idx = torch.topk(scores, k, largest=False)
# 通过集合差找到其余索引
rest_mask = torch.ones_like(scores, dtype=torch.bool)
rest_mask.scatter_(-1, min_idx, False)
rest_idx = all_idx[rest_mask]
或者直接用 torch.argsort(scores)[k:] 拿到除了最小 k 个之外的所有索引。这两种方法各有适用场景,前者在 k 很小时更快,后者代码更简洁,看个人偏好。
5.4 几个加速小妙招
最后分享几个我一直在用的加速小习惯:
- GPU 上优先用 topk:对大批量张量,
topk的并行度比sort好,深度模型训练中能明显减少等待时间。 - 避免频繁调用小 topk:如果循环里要反复对很小的张量取 topk,开销主要在 Python 调用和 kernel 启动上。可以考虑把循环改成 batch 形式,让一次 topk 处理完。
sorted=False可能更快:如果只需要选出来、不需要排序结果,把sorted设为False能省掉一部分内部排序工作。- 用 float16 或 bfloat16 时注意精度:半精度下
topk的排序稳定性可能受影响,对顺序敏感的任务建议回退到 float32 核对一次结果。
写到这里,“PyTorch中获取最小的k个元素”这件事基本覆盖完整了。从 API 参数到边界条件,从性能优化到实战场景,每一步都是我实际踩过的路。最后再提醒一句:遇到类似“找 TopK”的需求,先问自己三个问题——取最大还是最小、沿哪个维度、要不要索引。这三个问题想清楚,代码基本不会写错。
