如何为PyTorch张量生成对应的索引张量?
优化PyTorch多维索引张量的生成方法
直接用PyTorch内置的torch.indices()函数就能一步搞定,代码简洁且效率更高:
import torch t = torch.tensor([[0., 1., 2.], [3., 4., 5.]]) idx = torch.indices(t.shape).permute(1, 2, 0) print(idx)
输出结果和预期完全一致:
tensor([[[0, 0], [0, 1], [0, 2]], [[1, 0], [1, 1], [1, 2]]])
通用场景适配
对于任意形状为(d0, d1, ..., dn)的张量t,torch.indices(t.shape)会生成形状为(n+1, d0, d1, ..., dn)的张量,第一个维度对应各个轴的索引。通过permute(*range(1, t.ndim+1), 0)调整维度顺序后,就能得到形状为(d0, d1, ..., dn, n+1)的目标索引张量。
另外也可以用torch.meshgrid实现,写法稍复杂一点:
idx = torch.stack(torch.meshgrid(*[torch.arange(s) for s in t.shape], indexing='ij'), dim=-1)
这两种方法都避免了手动循环,完全依赖PyTorch的向量化操作,比原实现更高效、更简洁。
内容的提问来源于stack exchange,提问作者jmhummel
相关产品推荐
相关产品推荐

