【问题标题】:Update specific vector elements in PyTorch更新 PyTorch 中的特定向量元素
【发布时间】:2019-04-18 16:43:32
【问题描述】:

我有一个大向量要更新。我将通过向向量中的特定元素添加偏移量来更新它。我指定了一个要更新的索引向量(称为索引向量ix),对于每个索引,我指定一个要添加到该元素的值(称为值向量vals)。如果索引向量的所有条目都是唯一的,那么下面的代码就足够了:

vec = torch.zeros(4, dtype=torch.float)
ix = torch.tensor([0,2], dtype=torch.long)
vals = torch.tensor([0.2, 0.5], dtype=torch.float)
vec[ix] += vals

但是,如果ix 中有重复的索引,这将不起作用。对于重复索引的情况,一种简单的方法如下:

for i in range(len(ix)):
    vec[ix[i]] += vals[i]

但这不能很好地扩展 - 当ix 很大时它非常慢。有没有更快的方法来做到这一点?如果有一种快速的方法来汇总 vals 中在 ix 中具有相同索引的所有条目,那么解决方案应该很简单。

更新:
我找到了一种效果很好的解决方案,如下所述。我仍然希望得到更好的解决方案的反馈。

# get unique indices
ix_unique = torch.unique(ix)

# for each unique index, get sum of all vals with that index
vals_unique = torch.stack([
    torch.sum(torch.where(ix==i, vals, torch.zeros_like(vals))) 
    for i in ix_unique
])

# update vec
vec[ix_unique] += vals_unique

【问题讨论】:

  • 你可以写下自己的答案,向别人表明有解决办法!

标签: python numpy pytorch autograd


【解决方案1】:

对于您希望允许对同一个 ix 索引进行多次更新的情况,还有一个名为 pytorch_scatter 的库。 在这种情况下,例如然后 ix 中的 3 个相同索引将导致 3*val 被添加到该索引。

【讨论】:

    【解决方案2】:

    torch.index_add()

    import torch
    
    vec = torch.zeros(4, dtype=torch.float)
    ix = torch.tensor([0,0,2], dtype=torch.long)
    vals = torch.tensor([0.2,0.1,0.5], dtype=torch.float)
    torch.index_add(vec, 0, ix, vals)
    

    你会得到

    tensor([0.3000, 0.0000, 0.5000, 0.0000])
    

    参考:official doc

    【讨论】:

      猜你喜欢
      • 1970-01-01
      • 2022-08-09
      • 2021-05-25
      • 1970-01-01
      • 2021-03-09
      • 1970-01-01
      • 1970-01-01
      • 2011-02-07
      • 1970-01-01
      相关资源
      最近更新 更多