如何在PyTorch中提取高维张量指定维度的对角元素?
嘿,刚接触PyTorch遇到这种高维张量的索引问题太正常了,我给你几个高效的向量化解决方案,完全不用写循环,直接搞定👇
方法1:高级索引(最直观)
直接构造匹配的索引,利用PyTorch的广播机制一步提取对角元素:
import torch n = 2 # 替换成你的实际n值 m = 3 # 替换成你的实际m值 T = torch.randn(n, n, m, m) # 示例张量 # 构造索引:取0到n-1的整数,对应每个n×n块的对角位置 idx = torch.arange(n, device=T.device) # 直接索引,自动广播匹配维度 U = T[idx, idx, :, :] print(U.shape) # 输出 torch.Size([2, 3, 3]),完美符合需求!
原理很简单:idx是一维张量,PyTorch会自动把它广播成和后面:, :匹配的形状,刚好提取每个T[i,i,k,l]的元素,直接得到(n×m×m)的结果。
方法2:用torch.diagonal(高维支持)
你以为torch.diagonal只支持二维?其实它完全能处理高维张量!只需要指定要取对角的两个维度就行:
# 对第0和第1维度取对角(也就是每个n×n块的对角) diag_tensor = T.diagonal(dim1=0, dim2=1) # 调整维度顺序到我们需要的(n×m×m) U = diag_tensor.permute(2, 0, 1) print(U.shape) # 同样输出 torch.Size([2, 3, 3])
这里diagonal会把对角元素放在最后一个维度,所以用permute调整一下顺序就好,也是纯向量化操作。
方法3:爱因斯坦求和(优雅简洁)
如果你愿意尝试einsum的话,这是个非常直观的写法,一行搞定:
U = torch.einsum('iikl->ikl', T)
'iikl->ikl'的意思是:保留张量T中第一个i和第二个i相等的元素,然后按照i,k,l的顺序输出,刚好就是我们要的U[i,k,l] = T[i,i,k,l],形状自动匹配,可读性拉满!
这三个方法都完全避开了循环,效率拉满,随便选一个你觉得顺手的就行~
内容的提问来源于stack exchange,提问作者lionhmm
相关产品推荐
相关产品推荐

