【问题标题】:Max-pooling with complex masks in PyTorch在 PyTorch 中使用复杂掩码进行最大池化
【发布时间】: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]])

虽然可以,但运行速度很慢,看起来也不优雅! 有什么建议可以改进它以提高效率

另外,如果srcadj都附加了batches(batch, 5, 3)(batch, 5, 5)),如何让它工作?

【问题讨论】:

标签: pytorch


【解决方案1】:

我正在试验你的代码:

output = torch.zeros_like(src)
for index in range(adj.size(0)):
  nz = adj[index].nonzero().view(-1)
  output[index] = src.index_select(0, nz).max(0).values

瓶颈当然是for循环。首先想到的是使用某种分散函数。但是,这里的主要问题是邻居的数量可能因行而异。这意味着我们将无法在最大池化之前构建包含候选节点的张量。


一个可能的解决方案是创建一个类似于src 的辅助张量,其中第一个节点将包含占位符值(这些不应由最大池选择,即 em> 我们可以使用-inf)。我们可以使用包含索引的张量对该张量进行索引:与您的方法相比,我们将放置一个索引值为 0 的索引值,而不是使用 torch.nonzero() 删除零(参考第一个占位符行modified-src的位置)。

在实践中,它的样子如下:

对于辅助张量 src_,我将 -1s 作为占位符值。

>>> src_ = torch.cat((-torch.ones_like(src[:1]), src))
tensor([[-inf, -inf, -inf],
        [  0.,   1.,   2.],
        [  3.,   4.,   5.],
        [  6.,   7.,   8.],
        [  9.,  10.,  11.],
        [ 12.,  13.,  14.]])

我们可以将adj 矩阵转换为索引张量:

>>> index = torch.arange(1, adj.size(1) + 1)*adj
tensor([[1, 0, 3, 4, 0],
        [0, 2, 3, 4, 0],
        [1, 2, 0, 4, 5],
        [1, 2, 3, 0, 0],
        [0, 0, 3, 0, 5]])

为了更容易索引,我们将在第一个轴上展平index,索引src_,然后重新整形:

>>> indexed = src_[index.flatten(), :].reshape(*adj.shape, 3)
tensor([[[  0.,   1.,   2.],
         [-inf, -inf, -inf],
         [  6.,   7.,   8.],
         [  9.,  10.,  11.],
         [-inf, -inf, -inf]],

        ...

        [[-inf, -inf, -inf],
         [-inf, -inf, -inf],
         [  6.,   7.,   8.],
         [-inf, -inf, -inf],
         [ 12.,  13.,  14.]]])

终于可以max-pool了:

>>> indexed.max(dim=1).values
tensor([[ 9., 10., 11.],
        [ 9., 10., 11.],
        [12., 13., 14.],
        [ 6.,  7.,  8.],
        [12., 13., 14.]])

【讨论】:

    猜你喜欢
    • 1970-01-01
    • 1970-01-01
    • 2020-09-09
    • 2023-03-13
    • 2020-11-10
    • 1970-01-01
    • 2017-01-23
    • 1970-01-01
    • 2014-12-26
    相关资源
    最近更新 更多