【问题标题】:Any numpy/torch style to set value given an index ndarray and a flag ndarray?给定索引ndarray和标志ndarray的任何numpy/torch样式设置值?
【发布时间】:2019-10-10 09:36:54
【问题描述】:

我正在研究 PyTorch,目前我遇到了一个问题,我不知道如何以 Torch/numpy 风格解决它。例如,假设我有三个 PyTorch 张量

import torch
import numpy as np

indices = torch.from_numpy(np.array([[2, 1, 3, 0], [1, 0, 3, 2]]))
flags = torch.from_numpy(np.array([[False, False, False, True], [False, False, True, True]]))
tensor = torch.from_numpy(np.array([[2.8, 0.5, 1.2, 0.9], [3.1, 2.8, 1.3, 2.5]]))

这里的flags 是一个布尔标志张量,用于显示应该提取indices 中的哪些元素。鉴于提取的索引,我想将tensor 中的相应元素设置为指定的常量(例如 1e-30)。基于上面显示的示例,我想要

>>> sub_indices = indices.op1(flags)
>>> sub_indices
tensor([[0], [3, 2]])
>>> tensor.op2(sub_indices, 1e-30)
>>> tensor
tensor([[1e-30, 0.5, 1.2, 0.9], [3.1, 2.8, 1e-30, 1e-30]])

谁能帮忙解决?我正在使用列表理解,但我认为这种方式有点难看。我试过indices[flags],但它只返回一个一维数组[0, 3, 2],所以应用它会改变同一列0、2、3上的所有行

一些补充说明:

  • 无法确定flags 中每一行的“True”值数量
  • indices 的每一行都保证是序列0 ... N - 1 的排列

下面是示例代码的 numpy 版本,方便复制粘贴。我怀疑这是否可以以纯粹的 numpy 方式完成

import numpy as np

indices = np.array([[2, 1, 3, 0], [1, 0, 3, 2]])
flags = np.array([[False, False, False, True], [False, False, True, True]])
tensor = np.array([[2.8, 0.5, 1.2, 0.9], [3.1, 2.8, 1.3, 2.5]])

【问题讨论】:

    标签: python numpy pytorch


    【解决方案1】:

    您可以根据indicesflags 进行排序以创建mask,然后将mask 用作复用器。这是一个示例代码:

    indices = np.array([[2, 1, 3, 0], [1, 0, 3, 2]])
    flags = np.array([[False, False, False, True], [False, False, True, True]])
    tensor = np.array([[2.8, 0.5, 1.2, 0.9], [3.1, 2.8, 1.3, 2.5]])
    
    indices_sorted = indices.argsort(axis=1)
    mask = np.take_along_axis(flags, indices_sorted, axis=1)
    result = tensor * (1 - mask) + 1e-30 * mask
    

    我对 pytorch 不太熟悉,但我想收集一个参差不齐的张量并不是一个好主意。不过,即使在最坏的情况下,您也可以与 numpy 数组相互转换。

    【讨论】:

      【解决方案2】:

      @soloice 解决方案的 pytorch 版本。在 pytorch 中,使用torch.gather 代替torch.take

      indices = torch.tensor([[2, 1, 3, 0], [1, 0, 3, 2]])
      flags = torch.tensor([[False, False, False, True], [False, False, True, True]])
      tensor = torch.tensor([[2.8, 0.5, 1.2, 0.9], [3.1, 2.8, 1.3, 2.5]])
      
      indices_sorted = indices.argsort(axis=1)
      mask = torch.gather(flags, 1, indices_sorted).float()
      result = tensor * (1 - mask) + 1e-30 * mask
      

      【讨论】:

        猜你喜欢
        • 2018-10-10
        • 2018-02-16
        • 1970-01-01
        • 2023-03-13
        • 1970-01-01
        • 2021-12-15
        • 1970-01-01
        • 2015-09-14
        • 2021-08-06
        相关资源
        最近更新 更多