【发布时间】: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