Python多进程:如何让子进程复用父进程加载的只读大数据
问题场景
需要从磁盘读取一份大数据并执行只读操作,遇到以下问题:
- 使用
multiprocessing.Manager()或Array()实现跨进程共享时速度过慢 - 将大数据声明为全局变量后,每个子进程仍会重新从磁盘加载数据,耗时严重
当前内存充足,希望实现仅由父进程从磁盘加载一次数据,子进程直接复用内存中的副本,避免重复磁盘IO。
原始代码示例:
# main.py import argparse import numpy as np import multiprocessing as mp import time parser = argparse.ArgumentParser() parser.add_argument('-p', '--path', type=str) args = parser.parse_args() print('loading data from disk... may take a long time...') global_large_data = np.load(args.path) def worker(row_id): # 对global_large_data执行只读操作 time.sleep(0.01) print(row_id, np.sum(global_large_data[row_id])) def main(): pool = mp.Pool(mp.cpu_count()) pool.map(worker, range(global_large_data.shape[0])) pool.close() pool.join() if __name__ == '__main__': main()
执行命令:
$ python3 main.py -p /path/to/large_data.npy
解决方案
针对Unix-like系统(Linux/macOS)
Unix系统下multiprocessing默认用fork创建子进程,子进程会继承父进程的内存空间,且通过**写时复制(COW)**机制,只有修改数据时才会复制内存页。只读场景下子进程可直接复用父进程已加载的数据,无需重复读盘。
修正代码:
# main.py import argparse import numpy as np import multiprocessing as mp import time # 全局变量占位 global_large_data = None def worker(row_id): # 只读操作不会触发内存复制 time.sleep(0.01) print(row_id, np.sum(global_large_data[row_id])) def main(): parser = argparse.ArgumentParser() parser.add_argument('-p', '--path', type=str) args = parser.parse_args() global global_large_data print('loading data from disk... may take a long time...') # 父进程仅加载一次数据 global_large_data = np.load(args.path) # 创建进程池时,子进程继承已加载的内存数据 pool = mp.Pool(mp.cpu_count()) pool.map(worker, range(global_large_data.shape[0])) pool.close() pool.join() if __name__ == '__main__': main()
针对Windows系统
Windows用spawn方式创建进程,子进程会重新执行整个脚本,无法直接继承内存数据。此时用multiprocessing.shared_memory实现内存共享,避免重复加载:
兼容代码:
# main.py import argparse import numpy as np import multiprocessing as mp from multiprocessing import shared_memory import time def worker(row_id, shm_name, shape, dtype): # 连接父进程创建的共享内存 existing_shm = shared_memory.SharedMemory(name=shm_name) # 将共享内存映射为numpy数组 large_data = np.ndarray(shape, dtype=dtype, buffer=existing_shm.buf) # 只读操作 time.sleep(0.01) print(row_id, np.sum(large_data[row_id])) # 关闭共享内存连接(父进程负责销毁) existing_shm.close() def main(): parser = argparse.ArgumentParser() parser.add_argument('-p', '--path', type=str) args = parser.parse_args() print('loading data from disk... may take a long time...') large_data = np.load(args.path) # 创建共享内存块,存入数据 shm = shared_memory.SharedMemory(create=True, size=large_data.nbytes) shared_data = np.ndarray(large_data.shape, dtype=large_data.dtype, buffer=shm.buf) shared_data[:] = large_data[:] # 复制数据到共享内存 # 传递共享内存参数给子进程 pool = mp.Pool(mp.cpu_count()) pool.starmap(worker, [(row_id, shm.name, large_data.shape, large_data.dtype) for row_id in range(large_data.shape[0])]) pool.close() pool.join() # 销毁共享内存 shm.close() shm.unlink() if __name__ == '__main__': main()
执行命令
$ python3 main.py -p /path/to/large_data.npy
内容的提问来源于stack exchange,提问作者graphitump
相关产品推荐
相关产品推荐

