这是一种非常疯狂的方法,它涉及排序和索引,而不是添加新维度。这有点像np.unique 使用的基于排序的方法。
首先找到每一行的排序索引:
rows = np.repeat(np.arange(x.shape[0]), x.shape[1]) # [0, 0, 0, 1, 1, 1, 2, 2, 2, 3, 3, 3]
cols = np.argsort(x, axis=1).ravel() # [0, 2, 1, 2, 1, 0, 0, 1, 2, 1, 0, 2]
现在您可以为每列创建一个排序元素数组,包括未加权和加权。前者用于求和的指标,后者实际是求和的。
u = x[rows, cols] # [1, 1, 2, 1, 2, 3, 1, 2, 2, 1, 3, 3]
v = np.broadcast_to(w, x.shape)[rows, cols] # [0.3, 0.3, 0.4, 0.3, 0.4, 0.3, 0.3, 0.4, 0.3, 0.4, 0.3, 0.3]
您可以在以下位置找到要应用np.add.reduce 的断点:
row_breaks = np.diff(rows).astype(bool) # [0, 0, 1, 0, 0, 1, 0, 0, 1, 0, 0]
col_breaks = np.diff(u).astype(bool) # [0, 1, 1, 1, 1, 1, 1, 0, 1, 1, 0]
break_mask = row_breaks | col_breaks # [0, 1, 1, 1, 1, 1, 1, 0, 1, 1, 0]
breaks = np.r_[0, np.flatnonzero(break_mask) + 1] # [ 0, 2, 3, 4, 5, 6, 7, 9, 10]
现在您有了每行中相同数字的权重总和:
sums = np.add.reduceat(v, breaks) # [0.6, 0.4, 0.3, 0.4, 0.3, 0.3, 0.7, 0.4, 0.6]
但是您需要根据每行唯一元素的数量将它们分解为段:
unique_counts = np.add.reduceat(break_mask, np.arange(0, x.size, x.shape[1]))
unique_counts[-1] += 1 # The last segment will be missing from the mask: # [2, 3, 2, 2]
unique_rows = np.repeat(np.arange(x.shape[0]), unique_counts) # [0, 0, 1, 1, 1, 2, 2, 3, 3]
您现在可以对每个段进行排序以找到最大值:
indices = np.lexsort(np.stack((sums, unique_rows), axis=0)) # [1, 0, 2, 4, 3, 5, 6, 7, 8]
每次运行结束时的索引由下式给出:
max_inds = np.cumsum(unique_counts) - 1 # [1, 4, 6, 8]
所以最大总和是:
sums[indices[max_inds]] # [0.6, 0.4, 0.7, 0.6]
您可以解开索引内的索引以从每一行中获取正确的元素。请注意max_inds,以及依赖它的所有内容都与x.shape[1] 一样大小,正如预期的那样:
result = u[breaks[indices[max_ind]]]
这种方法看起来不是很漂亮,但它可能比在数组上使用额外维度更节省空间。此外,无论x 中的数字如何,它都能正常工作。请注意,我从未以任何方式减去任何内容或调整x。实际上,所有行都是独立处理的,构造breaks时,row_breaks打破了最大元素与下一个元素相同的重合。
TL;DR
享受:
def weighted_vote(x, w):
rows = np.repeat(np.arange(x.shape[0]), x.shape[1])
cols = np.argsort(x, axis=1).ravel()
u = x[rows, cols]
v = np.broadcast_to(w, x.shape)[rows, cols]
row_breaks = np.diff(rows).astype(bool)
col_breaks = np.diff(u).astype(bool)
break_mask = row_breaks | col_breaks
breaks = np.r_[0, np.flatnonzero(break_mask) + 1]
sums = np.add.reduceat(v, breaks)
unique_counts = np.add.reduceat(break_mask, np.arange(0, x.size, x.shape[1]))
unique_counts[-1] += 1
unique_rows = np.repeat(np.arange(x.shape[0]), unique_counts)
indices = np.lexsort(np.stack((sums, unique_rows), axis=0))
max_inds = np.cumsum(unique_counts) - 1
return u[breaks[indices[max_inds]]]
基准测试
基准测试以下列格式运行:
rows = ...
cols = ...
x = np.random.randint(cols, size=(rows, cols)) + 1
w = np.random.rand(cols)
%timeit weighted_vote_MP(x, w)
%timeit weighted_vote_JG(x, w)
assert (weighted_vote_MP(x, w) == weighted_vote_JG(x, w)).all()
我对@987654340@ 使用了以下概括,并进行了适当的更正:
def weighted_vote_JG(x, w):
i = np.arange(w.size) + 1
m = (x[None, ...] == i.reshape(-1, 1, 1))
return np.argmax(np.sum(m * w, axis=2), axis=0) + 1
行:100,列:10
MP: 440 µs ± 5.12 µs per loop (mean ± std. dev. of 7 runs, 1000 loops each)
* JG: 153 µs ± 796 ns per loop (mean ± std. dev. of 7 runs, 10000 loops each)
行:1000,列:10
MP: 2.53 ms ± 43.7 µs per loop (mean ± std. dev. of 7 runs, 100 loops each)
* JG: 1.03 ms ± 12 µs per loop (mean ± std. dev. of 7 runs, 1000 loops each)
行:10000,列:10
MP: 23.5 ms ± 200 µs per loop (mean ± std. dev. of 7 runs, 10 loops each)
* JG: 16.6 ms ± 67.4 µs per loop (mean ± std. dev. of 7 runs, 100 loops each)
行:100000,列:10
MP: 322 ms ± 3.11 ms per loop (mean ± std. dev. of 7 runs, 1 loop each)
* JG: 188 ms ± 858 µs per loop (mean ± std. dev. of 7 runs, 10 loops each)
行:100,列:100
* MP: 3.31 ms ± 257 µs per loop (mean ± std. dev. of 7 runs, 100 loops each)
JG: 12.6 ms ± 244 µs per loop (mean ± std. dev. of 7 runs, 100 loops each)
行:1000,列:100
* MP: 31 ms ± 159 µs per loop (mean ± std. dev. of 7 runs, 10 loops each)
JG: 134 ms ± 581 µs per loop (mean ± std. dev. of 7 runs, 10 loops each)
行:10000,列:100
* MP: 417 ms ± 7.06 ms per loop (mean ± std. dev. of 7 runs, 1 loop each)
JG: 1.42 s ± 126 ms per loop (mean ± std. dev. of 7 runs, 1 loop each)
行:100000,列:100
* MP: 4.94 s ± 25.9 ms per loop (mean ± std. dev. of 7 runs, 1 loop each)
JG: MemoryError: Unable to allocate 7.45 GiB for an array with shape (100, 100000, 100) and data type float64
故事的寓意:对于少量的列和权重,扩展解决方案更快。对于更多列,请改用我的版本。