如何利用torch.topk()索引提取每组前5长的2D向量坐标?
解决PyTorch中按组取前K个向量的维度匹配问题
这问题我之前也踩过坑!核心是PyTorch的索引广播规则在这儿给咱们挖了个小坑,先帮你理清楚为什么原来的代码会得到奇怪的形状,再给你两种靠谱的解决方法。
为什么原来的索引不对?
你拿到的tk.indices形状是[N,5],当直接用x[tk.indices]索引时,PyTorch会把这个索引当作替换第一个维度的取值,然后和原张量的[N,N_g,2]维度进行广播,最后就多出了一个不必要的N_g维度,导致结果变成[10,5,10,2],完全不符合咱们要的[N,5,2]。
两种正确实现方法
方法1:用torch.gather精准收集(推荐)
gather是PyTorch专门用来根据索引在指定维度上收集值的函数,完全适配咱们的场景:
import torch x = torch.randn(10, 10, 2) # N=10(batch),N_g=10(每组向量数) x_len = (x**2).sum(dim=2).sqrt() tk = x_len.topk(5) # 给indices增加最后一个维度,再repeat适配向量的2D维度 index = tk.indices.unsqueeze(-1).repeat(1, 1, 2) x_top5 = torch.gather(x, dim=1, index=index) print(x_top5.shape) # 输出: torch.Size([10, 5, 2])
- 解释:
index需要和输入张量x的维度数一致(3维),所以先给tk.indices加一个最后一维变成[10,5,1],再repeat这个维度两次,变成[10,5,2],这样gather在dim=1(也就是每组向量的维度)上收集时,就能准确取到每个batch里前5长的2D向量。
方法2:手动构建batch索引匹配维度
如果觉得gather有点绕,也可以手动构建batch的索引,让索引维度完全匹配:
import torch x = torch.randn(10, 10, 2) x_len = (x**2).sum(dim=2).sqrt() tk = x_len.topk(5) # 创建batch维度的索引,形状[10,1],和tk.indices[10,5]广播成[10,5] batch_idx = torch.arange(x.size(0)).unsqueeze(1) # 用batch_idx和tk.indices一起索引,精准定位每个batch里的目标向量 x_top5 = x[batch_idx, tk.indices] print(x_top5.shape) # 输出: torch.Size([10, 5, 2])
- 解释:
batch_idx每个元素对应一个batch的序号,和tk.indices结合后,相当于告诉PyTorch:“在第i个batch里,取第tk.indices[i,j]个向量”,这样索引出来的结果维度就完全符合预期了。
内容的提问来源于stack exchange,提问作者ihdv
相关产品推荐
相关产品推荐

