如何在PyTorch中高效构造符合Q[i,j,k]=P[A[i,j],k]的三维张量?
高效构造张量Q的PyTorch实现方案
我们可以通过PyTorch的高级索引或torch.gather两种高效方式实现需求,以下是具体方案:
方法一:高级索引(简洁直观)
利用张量的广播特性直接索引,代码最简洁:
# 假设A和P均为已定义的torch张量 Q = P[A[:, :, None], :]
解释:
A[:, :, None]将二维张量A的形状从(n,n)扩展为(n,n,1)- 用该三维张量索引二维矩阵P时,PyTorch会自动广播维度,最终生成形状为
(n,n,n)的Q,恰好满足Q[i,j,k] = P[A[i,j],k]
方法二:使用torch.gather
通过调整张量维度适配gather的使用逻辑:
# 将P扩展为三维张量,形状变为(n, 1, n) P_expanded = P.unsqueeze(1) # 把A扩展为(n,n,n)的索引张量,对应P的行索引 A_index = A.unsqueeze(-1).repeat(1, 1, n) # 在第0维上根据索引取值 Q = torch.gather(P_expanded, dim=0, index=A_index)
解释:
A_index的每个位置存储了对应P的行索引,与P的列维度对齐gather在第0维度上,按照A_index的指引提取P的对应行元素,最终得到目标张量Q
说明
两种方法均基于PyTorch底层优化,性能相近,高级索引的写法更简洁,推荐优先使用。可通过torch.allclose(Q1, Q2)验证两种方法的输出结果一致。
内容的提问来源于stack exchange,提问作者Vezen BU
相关产品推荐
相关产品推荐

