【问题标题】:Numpy random choice, replacement only along one axisNumpy随机选择,仅沿一个轴替换
【发布时间】:2018-07-12 15:55:10
【问题描述】:

我需要从一个数组中采样一堆点对。我希望每对包含两个 DISTINCT 点,但这些点可能在不同对之间重复。

例如,如果我的数组是X=np.array([1,1,2,3]),那么:

>>> sample_pairs(X, n=4)
... [[1,1], [2,3], [1,2], [1,3]] # this is fine
>>> sample_pairs(X, n=4)
... [[1,1], [2,2], [3,3], [1,3]] # this is not okay

有没有一种好方法可以将其作为矢量化操作来完成?

【问题讨论】:

  • 实际用例是我试图引导成对距离的分布,而不计算所有成对距离,即O(n^2)
  • 向我们展示您打算如何使用这些配对。甚至可能是没有修剪的样本计算。通常在numpy 中,在一个向量化操作中完成所有计算会更快,而不是花费额外的时间来跳过冗余计算。
  • 我打算对一堆随机对进行采样,计算这些对之间的距离,然后返回结果距离列表的平均值和标准差。请注意,我可以计算成对距离,但我希望我能在线性时间内得到成对距离的均值和标准差的合理近似值。
  • 这个问题的目标是从X中随机选择点对
  • 用替换来计算整个样本似乎更简单,然后如果你不想要那些相同的点对就扔掉。

标签: python numpy random


【解决方案1】:

要对一对没有替换的样品进行采样,您可以使用np.random.choice

np.random.choice(X, size=2, replace=False)

或者,要一次采样多个元素,请注意,所有可能的对都可以由range(len(X)*(len(X)-1)/2) 的元素表示,并使用np.random.randint 从中采样。

combs = np.array(list(itertools.combinations(X, 2)))
sample = np.random.randint(len(combs), size=10)
combs[sample[np.newaxis]]

跟进@user2357112 的评论,根据 OP 自己的回答,他们似乎并不关心样本大小本身是否是确定性的,并指出使用 Mersenne Twister 进行采样比基本算术运算慢,如果是不同的解决方案X 太大了,生成组合是不可行的

sample = np.random.randint(len(X)**2, size=N)
i1 = sample // len(X)
i2 = sample % len(X)
X[np.vstack((i1, i2)).T[i1 != i2]]

这会产生一个平均大小为N * (1 - 1/len(X))的样本。

【讨论】:

  • 这仅对一对进行采样。我可以采样n 对然后.vstack 它们,但这似乎效率低下。我希望一个向量操作就能让我得到整个shebang。
  • 还不错;也添加了一种方法。
  • 我投了赞成票,但想知道他们是否是比combinations -> list 获得多对的更好方法,因为对于大数组来说,这很快就会失控
  • list(itertools.combinations(X,2))O(n^2),比.vstack 选项效率还要低....
  • @Scott 计算复杂性并不是一切;常量在矢量化操作中扮演着重要角色,因此如果这还不够好,您需要提供更接近实际用例的示例。
【解决方案2】:

这是@user2357112 的解决方案:

def sample_indices(X, n=4):
    pair_indices = np.random.randint(X.shape[0]**2, size=n)
    pair_indices = np.hstack(((pair_indices // X.shape[0]).reshape((-1,1)), (pair_indices % X.shape[0]).reshape((-1,1))))
    good_indices = pair_indices[:,0] != pair_indices[:,1]
    return X[pair_indices[good_indices]]

【讨论】:

  • np.random.randint(X.shape[0], size=(n, 2)) 可能会明显更快,并避免在您似乎担心的维度中进行分配。
  • 还请注意,这实际上并没有提供大小为n 的样本,因为你扔掉了其中的sqrt(len(X)) om 平均值。如果这不是问题,更快的解决方案是生成样本 s = np.random.randint(len(X)**2, size=n) 并使用 s // len(X)s % len(X) 提供索引(因为这些简单的操作比运行 Mersenne Twister 的额外轮次要快得多,所以加速大约翻了一番)。
  • 我将更新此解决方案以反映您的出色建议。
  • 在哪里? (但是是的,sqrt(n) 至少应该是 1/len(X)。)
  • “在哪里?”... NM。我很困惑。 :)
猜你喜欢
  • 1970-01-01
  • 2021-07-20
  • 2019-04-12
  • 1970-01-01
  • 1970-01-01
  • 2017-09-27
  • 1970-01-01
  • 1970-01-01
  • 1970-01-01
相关资源
最近更新 更多