【发布时间】:2021-01-09 08:16:03
【问题描述】:
假设我有一个形状为(5, 3) 的矩阵src 和一个形状为(5, 5) 的布尔矩阵(5, 5),如下所示,
src = tensor([[ 0, 1, 2],
[ 3, 4, 5],
[ 6, 7, 8],
[ 9, 10, 11],
[12, 13, 14]])
和
adj = tensor([[1, 0, 1, 1, 0],
[0, 1, 1, 1, 0],
[1, 1, 0, 1, 1],
[1, 1, 1, 0, 0],
[0, 0, 1, 0, 1]])
我们可以将src中的每一行作为一个节点嵌入,将adj中的每一行作为哪些节点是邻域的指标。
我的目标是在src 中的每个节点的所有邻域节点嵌入中运行一个最大池。
例如,由于第 0 个节点的邻域节点(包括其自身)是 0, 2, 3,因此我们在 [0, 1, 2]、[6, 7, 8]、[ 9, 10, 11] 上计算最大池化,并将更新后的嵌入 [ 9, 10, 11] 引导到更新src_update中的第0个节点。
我写的一个简单的解决方案是
src_update = torch.zeros_like(src)
for index in range(adj.size(0)):
list_of_non_zero = adj[index].nonzero().view(-1)
mat_non_zero = torch.index_select(src, 0, list_of_non_zero)
src_update[index] = torch.sum(mat_non_zero, dim=0)
而src_update 更新为:
tensor([[ 9, 10, 11],
[ 9, 10, 11],
[12, 13, 14],
[ 6, 7, 8],
[12, 13, 14]])
虽然可以,但运行速度很慢,看起来也不优雅! 有什么建议可以改进它以提高效率?
另外,如果src和adj都附加了batches((batch, 5, 3),(batch, 5, 5)),如何让它工作?
【问题讨论】:
-
建议你看一下pytorch scatter库:pytorch-scatter.readthedocs.io/en/1.3.0/functions/max.html
标签: pytorch