【问题标题】:Indexing a multi-dimensional tensor with a tensor in PyTorch在 PyTorch 中使用张量索引多维张量
【发布时间】:2019-02-05 02:46:21
【问题描述】:

我有以下代码:

a = torch.randint(0,10,[3,3,3,3])
b = torch.LongTensor([1,1,1,1])

我有一个多维索引b 并想用它来选择a 中的单个单元格。如果 b 不是张量,我可以这样做:

a[1,1,1,1]

返回正确的单元格,但是:

a[b]

不起作用,因为它只选择了四次a[1]

我该怎么做?谢谢

【问题讨论】:

    标签: pytorch tensor


    【解决方案1】:

    您可以使用chunkb 拆分为4 个,然后使用分块的b 来索引您想要的特定元素:

    >> a = torch.arange(3*3*3*3).view(3,3,3,3)
    >> b = torch.LongTensor([[1,1,1,1], [2,2,2,2], [0, 0, 0, 0]]).t()
    >> a[b.chunk(chunks=4, dim=0)]   # here's the trick!
    Out[24]: tensor([[40, 80,  0]])
    

    它的好处是它可以很容易地推广到a的任何维度,你只需要将卡盘的数量与a的维度相等。

    【讨论】:

    • 能够同时使用多个索引的额外奖励,我在我的问题中没有考虑到这一点。对此进行了测试,它可以工作,尽管值得注意的是我需要压缩输出。谢谢!
    • @Chum-ChumScarecrows 感谢您的接受,但 AFAIK dennlinger's answer 也推广到多个索引。我想你应该接受他的。
    【解决方案2】:

    一个更优雅(和更简单)的解决方案可能是将b 简单地转换为一个元组:

    a[tuple(b)]
    Out[10]: tensor(5.)
    

    我很想知道它如何与“常规”numpy 一起工作,并找到了一篇相关文章很好地解释了这一点 here

    【讨论】:

    • 有没有办法让这个解决方案与索引列表一起工作?
    • 原来a[list(b)] 也有效。有趣的。还是您指的是“列表中的元素列表”(即,类似于b = [[1,1,1,1], [1,1,1,2], [2,3,1,2]]
    • 嗯...我们可以在不将索引张量转换为元组的情况下做到这一点吗? (假设它很大并且驻留在 GPU 上,制作一个元组会将所有值拉到 CPU 上,这既是开销又迫使 GPU 在 CPU 上等待,反之亦然)。
    • 我已经有一段时间没有使用它了,所以我不能自信地回答你的问题。我的直觉告诉我这是不可能的,你将不得不移动数据。不过,我很高兴被证明是错误的,所以也许这可能是一个单独的问题?
    猜你喜欢
    • 2020-09-26
    • 2019-09-15
    • 1970-01-01
    • 2020-03-28
    • 2019-11-26
    • 2021-11-12
    • 1970-01-01
    • 2019-11-06
    • 2021-12-16
    相关资源
    最近更新 更多