You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

如何在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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.08.18 06:50:29