【问题标题】:Stacking tensors in a list of tuples of tensors在张量元组列表中堆叠张量
【发布时间】:2019-12-03 02:13:45
【问题描述】:

我有一个 PyTorch 张量元组的列表。它看起来像这样:

[
    (tensor([1, 2, 3]),  tensor([4, 5, 6, 7]),  tensor([8])),
    (tensor([9, 10,11]), tensor([11,12,13,14]), tensor([15])),
    (tensor([16,17,18]), tensor([19,20,21,22]), tensor([23])),
    ...
]

每列中的张量(即,位于其各自元组的 k 位置的张量)共享相同的形状。我想在每列中堆叠张量,以便最终得到一个元组,每个值都是沿列维度连接的张量。

在这种情况下,输出元组将具有三个值,如下所示:

(
 tensor([[1,2,3], [9,10,11], [16,17,18]]),

 tensor([[4,5,6,7], [11,12,13,14], [19,20,21,22]],

 tensor([[8],[15],[23])
)

这是一个虚构的例子。我想对任意长度的元组和任意大小的张量执行此操作。使用 PyTorch 快速执行此类连接的最佳方法是什么?

【问题讨论】:

    标签: python pytorch torch


    【解决方案1】:

    如果有人让自己陷入同样复杂的场景,我可以用一个可爱的单线来解决它:

    tuple(map(torch.stack, zip(*x)))
    

    在这种情况下,x 是我上面提到的原始列表。这行代码将x 转换为所需的确切格式。

    【讨论】:

    • 我认为您的意思是 torch.stack 而不是 torch.tensor
    • 很好的解决方案,谢谢。顺便说一句,我需要将张量堆叠在第一个暗淡上,因为 torch.cat 提供了解决方案
    猜你喜欢
    • 1970-01-01
    • 1970-01-01
    • 2020-08-30
    • 2018-06-05
    • 2019-04-19
    • 2020-03-05
    • 1970-01-01
    • 1970-01-01
    • 1970-01-01
    相关资源
    最近更新 更多