f. *_* c. 5 python performance numpy
我想从 numpy.searchsorted() 的结果中生成一个掩码:
import numpy as np
# generate test examples
x = np.random.rand(1000000)
y = np.random.rand(200)
# sort x
idx = np.argsort(x)
sorted_x = np.take_along_axis(x, idx, axis=-1)
# searchsort y in x
pt = np.searchsorted(sorted_x, y)
Run Code Online (Sandbox Code Playgroud)
pt是一个数组。然后我想创建一个大小(200, 1000000)为 True 值的布尔掩码,当它的索引为 时idx[0:pt[i]],我想出了一个像这样的 for 循环:
mask = np.zeros((200, 1000000), dtype='bool')
for i in range(200):
mask[i, idx[0:pt[i]]] = True
Run Code Online (Sandbox Code Playgroud)
任何人都有加速for循环的想法?
根据新发现的OP\'s comments状态仅y实时变化的信息,我们可以预处理很多东西x,从而做得更好。我们将创建一个哈希数组来存储阶梯掩码。对于涉及 的部分y,我们将简单地使用获得的索引对哈希数组进行索引searchsorted ,这将近似最终的掩码数组。鉴于 numba 的参差不齐的性质,分配剩余 bool 的最后一步可以卸载到 numba。如果我们决定扩大 的长度,这也应该是有益的y。
让我们看一下实现。
\n预处理x:
sidx = x.argsort()\nssidx = x.argsort().argsort()\n\n# Choose a scale factor. \n# 1. A small one would store more mapping info, hence faster but occupy more mem\n# 2. A big one would store less mapping info, hence slower, but memory efficient.\nscale_factor = 100\nmapar = np.arange(0,len(x),scale_factor)[:,None] > ssidx\nRun Code Online (Sandbox Code Playgroud)\n剩余步骤y:
import numba as nb\n\n@nb.njit(parallel=True,fastmath=True)\ndef array_masking3(out, starts, idx, sidx):\n N = len(out)\n for i in nb.prange(N):\n for j in nb.prange(starts[i], idx[i]):\n out[i,sidx[j]] = True\n return out\n\nidx = np.searchsorted(x,y,sorter=sidx)\ns0 = idx//scale_factor\nstarts = s0*scale_factor\nout = mapar[s0]\nout = array_masking3(out, starts, idx, sidx)\nRun Code Online (Sandbox Code Playgroud)\n标杆管理
\nIn [2]: x = np.random.rand(1000000)\n ...: y = np.random.rand(200)\n\nIn [3]: ## Pre-processing step with "x"\n ...: sidx = x.argsort()\n ...: ssidx = x.argsort().argsort()\n ...: scale_factor = 100\n ...: mapar = np.arange(0,len(x),scale_factor)[:,None] > ssidx\n\nIn [4]: %%timeit\n ...: idx = np.searchsorted(x,y,sorter=sidx)\n ...: s0 = idx//scale_factor\n ...: starts = s0*scale_factor\n ...: out = mapar[s0]\n ...: out = array_masking3(out, starts, idx, sidx)\n41 ms \xc2\xb1 141 \xc2\xb5s per loop (mean \xc2\xb1 std. dev. of 7 runs, 10 loops each)\n\n # A 1/10th smaller hashing array has similar timings\nIn [7]: scale_factor = 1000\n ...: mapar = np.arange(0,len(x),scale_factor)[:,None] > ssidx\n\nIn [8]: %%timeit\n ...: idx = np.searchsorted(x,y,sorter=sidx)\n ...: s0 = idx//scale_factor\n ...: starts = s0*scale_factor\n ...: out = mapar[s0]\n ...: out = array_masking3(out, starts, idx, sidx)\n40.6 ms \xc2\xb1 196 \xc2\xb5s per loop (mean \xc2\xb1 std. dev. of 7 runs, 10 loops each)\n\n# @silgon\'s soln \nIn [5]: %timeit x[np.newaxis,:] < y[:,np.newaxis]\n138 ms \xc2\xb1 896 \xc2\xb5s per loop (mean \xc2\xb1 std. dev. of 7 runs, 10 loops each)\nRun Code Online (Sandbox Code Playgroud)\n这借用了很好的一部分OP\'s solution。
import numba as nb\n\n@nb.njit(parallel=True)\ndef array_masking2(mask1D, mask_out, idx, pt):\n n = len(idx)\n for j in nb.prange(len(pt)):\n if mask1D[j]:\n for i in nb.prange(pt[j],n):\n mask_out[j, idx[i]] = False\n else:\n for i in nb.prange(pt[j]):\n mask_out[j, idx[i]] = True\n return mask_out\n\ndef app2(idx, pt):\n m,n = len(pt), len(idx) \n mask1 = pt>len(x)//2\n mask2 = np.broadcast_to(mask1[:,None], (m,n)).copy()\n return array_masking2(mask1, mask2, idx, pt)\nRun Code Online (Sandbox Code Playgroud)\n所以,我们的想法是,一旦我们有超过一半的索引需要设置,我们在将这些行预先分配为 all 后True切换到 set 。这会减少内存访问,从而显着提高性能。FalseTrue
OP的解决方案:
\n@nb.njit(parallel=True,fastmath=True)\ndef array_masking(mask, idx, pt):\n for j in nb.prange(pt.shape[0]):\n for i in nb.prange(pt[j]):\n mask[j, idx[i]] = True\n return mask\n\ndef app1(idx, pt):\n m,n = len(pt), len(idx) \n mask = np.zeros((m, n), dtype=\'bool\')\n return array_masking(mask, idx, pt)\nRun Code Online (Sandbox Code Playgroud)\n时间安排 -
\nIn [5]: np.random.seed(0)\n ...: x = np.random.rand(1000000)\n ...: y = np.random.rand(200)\n\nIn [6]: %timeit app1(idx, pt)\n264 ms \xc2\xb1 8.91 ms per loop (mean \xc2\xb1 std. dev. of 7 runs, 1 loop each)\n\nIn [7]: %timeit app2(idx, pt)\n165 ms \xc2\xb1 3.43 ms per loop (mean \xc2\xb1 std. dev. of 7 runs, 10 loops each)\nRun Code Online (Sandbox Code Playgroud)\n
| 归档时间: |
|
| 查看次数: |
445 次 |
| 最近记录: |