【问题标题】:Speeding up tensor concatenation加速张量连接
【发布时间】:2020-07-30 21:15:43
【问题描述】:

我需要连接一长串小张量。每个小张量都是给定(非常简单)常数矩阵的切片。代码如下:

max_node, counter = 0, 0
batch_size, n_days = (1000, 10)
n_interactions_in = torch.randint(low=100,high=200,size=(batch_size,n_days), dtype=torch.long)
max_interactions = n_interactions_in.max()
delay_table = torch.arange(n_days, device=device, dtype=torch.float).expand([max_interactions, n_days]).t().contiguous()
delay_table = n_days - delay_table - 1
edge_delay_buf = []
for b in range(batch_size):
     delay_vec = [delay_table[d, :n_interactions_in[b, d]] for d in range(n_days)]
     edge_delay_buf.append(torch.cat(delay_vec))
res = torch.cat(edge_delay_buf)

这需要很多时间。有没有办法有效地并行化 edge_delay_buf 中每个元素的创建? 我尝试了多种变体,例如用列表连接替换 for 循环,其中结果是列表列表,然后展平列表并在展平列表上应用 torch.cat。然而,并没有太大改善。由于某种原因,切片操作耗时过长。

有没有办法让切片更快?有没有办法让循环更高效/并行?

注意:虽然我在这个例子中使用了 torch,但我也可以使用 numpy. 注意 2:对于在其他论坛中重复发布的帖子,我深表歉意。

【问题讨论】:

  • numpy 切片不需要很长时间;它使view。但最终,当连接所有这些视图时,它必须将所有值复制到新数组中。显然,如果batch_size 很大,那么该列表追加步骤将需要时间,否则我怀疑这是torch.cat 步骤是最重要的消费者。但是您应该能够对这些步骤进行时间测试。

标签: python numpy pytorch concatenation


【解决方案1】:

先用列表追加替换内部串联,最后只做一次串联,应该会快得多。

max_node, counter = 0, 0
batch_size, n_days = (1000, 10)
n_interactions_in = torch.randint(low=100,high=200,size=(batch_size,n_days), dtype=torch.long)
max_interactions = n_interactions_in.max()
delay_table = torch.arange(n_days, device=device, dtype=torch.float).expand([max_interactions, n_days]).t().contiguous()
delay_table = n_days - delay_table - 1
edge_delay_buf = []
for b in range(batch_size):
     delay_vec = [delay_table[d, :n_interactions_in[b, d]] for d in range(n_days)]
     edge_delay_buf += delay_vec
res = torch.cat(edge_delay_buf)

然后,如果它仍然不够快,可以通过一次提取所有索引来提高效率。让我们看看你有一个形状为[N,M]的矩阵A,你实际上可以提取一些元素做A [B,C],其中B是长度为K的向量,C是形状为[L,K的矩阵]。也许它可以满足您的需求。

【讨论】:

    猜你喜欢
    • 2022-08-18
    • 2019-06-10
    • 2021-05-17
    • 2021-11-20
    • 2018-09-20
    • 2021-07-03
    • 2017-08-20
    • 2019-07-10
    • 2019-05-06
    相关资源
    最近更新 更多