如何在NumPy二维数组的指定索引位置实现最快写入?
如何在二维NumPy数组中通过索引数组快速写入数据?
针对你提到的大型布尔型二维NumPy数组写入场景,以下是几种比现有方案更快的实现方式,尤其是基于Numba的优化方案:
一、利用Numba JIT编译加速
Numba可将Python函数编译为机器码,大幅降低循环开销,对这类逐元素索引写入的场景优化效果显著。
1. 通用场景(row_indexes为任意数组)
适用于row_indexes非连续序列的情况:
import numba as nb @nb.njit(nb.void(nb.boolean[:,:], nb.uint32[:], nb.uint32[:], nb.boolean[:]), fastmath=True) def numba_write(buffer, row_idx, col_idx, data): n = row_idx.shape[0] for i in range(n): buffer[row_idx[i], col_idx[i]] = data[i]
性能测试(首次运行含编译开销,后续调用无额外成本):
%timeit numba_write(buffer, row_indexes, w_idx, data) # 典型结果:~0.8 ms ± 10 µs per loop(比np.put快约1倍)
2. 优化场景(row_indexes为连续序列)
由于你的row_indexes是np.arange(n_rows)的连续行索引,可省略对row_indexes的数组访问,进一步提速:
@nb.njit(nb.void(nb.boolean[:,:], nb.uint32[:], nb.boolean[:]), fastmath=True) def numba_write_continuous(buffer, col_idx, data): n = col_idx.shape[0] for i in range(n): buffer[i, col_idx[i]] = data[i]
性能测试:
%timeit numba_write_continuous(buffer, w_idx, data) # 典型结果:~0.45 ms ± 5 µs per loop(比原最快方案快约3倍)
二、扁平化索引直接赋值(无额外依赖)
若不想引入Numba,可直接计算扁平化索引后赋值,性能略优于np.put:
flat_idx = row_indexes * n_cols + w_idx buffer.ravel()[flat_idx] = data
性能测试:
%timeit buffer.ravel()[row_indexes * n_cols + w_idx] = data # 典型结果:~1.6 ms ± 15 µs per loop
关键优化说明
- Numba编译开销:首次调用Numba函数会产生几十毫秒的编译时间,但后续重复调用无额外成本,适合需要多次执行写入操作的场景。
- 内存顺序:你的buffer已设置为
order="C"(行优先),与Numba、NumPy的默认内存访问模式一致,能最大化缓存命中率。 - 布尔数组特性:布尔数组仅占用1字节/元素,内存访问效率更高,上述优化方案可充分利用这一特性。
内容的提问来源于stack exchange,提问作者Wang Lee
相关产品推荐
相关产品推荐

