如何在Python多进程间共享numpy结构化数组?
在多进程中共享NumPy结构化数组的可行方案
我来帮你梳理下在多进程中共享NumPy结构化数组的几种可行方案,结合你的需求(包括Cython的使用)逐一说明:
1. 用multiprocessing.Array包装共享内存(兼容所有Python版本)
NumPy结构化数组的内存是连续的,所以我们可以把它的数据拷贝到mp.Array创建的共享内存缓冲区里,然后在子进程中重新构建数组。这种方法不需要依赖高版本Python,兼容性好。
修改你的代码如下:
import numpy as np import multiprocessing as mp def modi(arg): i, shared_buf, dtype = arg # 从共享内存缓冲区重建结构化数组 shared_arr = np.frombuffer(shared_buf.get_obj(), dtype=dtype) # 这里你原代码里的mdistance应该是笔误,改成了distance if shared_arr['distance'][i][0] == 4: return i shared_arr['distance'][i] += 2 if __name__ == '__main__': eetype = [('coordinate', 'f8', (2,)), ('file_id', '<U11'), ('distance', 'f8', (2,))] aa = np.zeros(5, dtype=eetype) aa['file_id'] = np.array(['aa', 'bb', 'cc', 'dd', 'ee']) aa['coordinate'] = np.array([[1,1],[2,2],[3,3],[4,4],[5,5]]) aa['distance'] = np.array([[1,1],[2,2],[3,3],[4,4],[5,5]]) print("初始数组:") print(aa) # 创建共享内存缓冲区,大小等于数组总字节数 # lock=False表示不自动加锁,如果你需要多进程同步修改,建议保留默认的lock=True或者自己加锁 shared_buf = mp.Array('b', aa.nbytes, lock=False) # 将原数组数据拷贝到共享内存 np.frombuffer(shared_buf.get_obj(), dtype=aa.dtype)[:] = aa[:] with mp.Pool(4) as p: args = [(i, shared_buf, aa.dtype) for i in range(len(aa))] results = list(p.map(modi, args)) # 从共享内存读取更新后的数组 updated_aa = np.frombuffer(shared_buf.get_obj(), dtype=aa.dtype) print("\n更新后的数组:") print(updated_aa)
注意:如果多个进程可能同时修改同一个数组元素,一定要加锁(比如把lock=False改成lock=True,或者在子进程里用mp.Lock),否则会出现数据竞争导致结果错误。
2. 使用multiprocessing.shared_memory(Python 3.8+ 推荐)
这是Python 3.8引入的更简洁的共享内存API,专门为NumPy数组这类需要共享内存的场景设计,不需要手动处理字节拷贝:
import numpy as np import multiprocessing as mp from multiprocessing import shared_memory def modi(arg): i, shm_name, dtype, shape = arg # 连接到主进程创建的共享内存块 existing_shm = shared_memory.SharedMemory(name=shm_name) # 重建结构化数组 shared_arr = np.ndarray(shape, dtype=dtype, buffer=existing_shm.buf) if shared_arr['distance'][i][0] == 4: return i shared_arr['distance'][i] += 2 # 关闭共享内存连接(不要unlink,主进程负责清理) existing_shm.close() if __name__ == '__main__': eetype = [('coordinate', 'f8', (2,)), ('file_id', '<U11'), ('distance', 'f8', (2,))] aa = np.zeros(5, dtype=eetype) aa['file_id'] = np.array(['aa', 'bb', 'cc', 'dd', 'ee']) aa['coordinate'] = np.array([[1,1],[2,2],[3,3],[4,4],[5,5]]) aa['distance'] = np.array([[1,1],[2,2],[3,3],[4,4],[5,5]]) print("初始数组:") print(aa) # 创建共享内存块,大小等于数组总字节数 shm = shared_memory.SharedMemory(create=True, size=aa.nbytes) # 创建指向共享内存的NumPy数组 shared_aa = np.ndarray(aa.shape, dtype=aa.dtype, buffer=shm.buf) # 拷贝数据到共享数组 shared_aa[:] = aa[:] with mp.Pool(4) as p: args = [(i, shm.name, aa.dtype, aa.shape) for i in range(len(aa))] results = list(p.map(modi, args)) print("\n更新后的数组:") print(shared_aa) # 清理共享内存(必须主进程执行) shm.close() shm.unlink()
这种方法代码更简洁,也更符合现代Python的用法,是优先推荐的方案。
3. 传递指针结合Cython(高性能需求)
你提到想通过传递指针实现共享,这在Cython里完全可行,而且性能会很高,因为直接操作底层内存。需要注意的是,必须保证主进程的数组在子进程运行期间不被销毁,同时要处理好同步。
首先写Cython的工作函数(比如worker.pyx):
import numpy as np cimport numpy as np from multiprocessing import Lock # 定义和NumPy结构化数组完全匹配的C结构体 ctypedef struct EEType: double coordinate[2] wchar_t file_id[11] # 对应Python里的<U11(注意:不同系统wchar_t大小可能不同,需要匹配) double distance[2] cdef void modify_element(EEType* arr, int idx, Lock lock): # 加锁避免多进程竞争 with lock: if arr[idx].distance[0] == 4: return arr[idx].distance[0] += 2 arr[idx].distance[1] += 2 def modi(arg): cdef int i = arg[0] # 将传递的整数转回指针 cdef EEType* ptr = <EEType*>arg[1] cdef Lock lock = arg[2] modify_element(ptr, i, lock) return i
然后在主进程中,获取数组的底层指针(转成整数传递,因为指针不能直接序列化),并调用子进程:
import numpy as np import multiprocessing as mp # 假设Cython编译后的模块名为worker from worker import modi if __name__ == '__main__': eetype = [('coordinate', 'f8', (2,)), ('file_id', '<U11'), ('distance', 'f8', (2,))] aa = np.zeros(5, dtype=eetype) aa['file_id'] = np.array(['aa', 'bb', 'cc', 'dd', 'ee']) aa['coordinate'] = np.array([[1,1],[2,2],[3,3],[4,4],[5,5]]) aa['distance'] = np.array([[1,1],[2,2],[3,3],[4,4],[5,5]]) print("初始数组:") print(aa) # 获取数组底层数据的指针,转成整数传递 arr_ptr = aa.__array_interface__['data'][0] # 创建锁用于同步 lock = mp.Lock() with mp.Pool(4) as p: args = [(i, arr_ptr, lock) for i in range(len(aa))] results = list(p.map(modi, args)) print("\n更新后的数组:") print(aa)
关键注意点:
- Cython里的结构体必须和NumPy的dtype完全匹配(字段顺序、类型、大小),否则会出现内存访问错误。
- 指针转成整数传递是因为Python不能直接序列化指针对象。
- 必须确保主进程的
aa数组在所有子进程结束前不被销毁,否则会出现悬空指针。 - 多进程修改时一定要加锁,避免数据损坏。
关于Ctype适配的补充
你提到觉得没有合适的Ctype适配结构化数组,其实可以手动创建对应的ctypes结构体:
import ctypes class EEType(ctypes.Structure): _fields_ = [ ('coordinate', ctypes.c_double * 2), ('file_id', ctypes.c_wchar * 11), # 对应<U11类型 ('distance', ctypes.c_double * 2) ] # 将NumPy数组转换为ctypes数组 ctypes_arr = (EEType * len(aa)).from_buffer(aa)
不过这种方法需要手动匹配每个字段的类型,比较繁琐,不如前面的共享内存方法直接,适合需要和C代码交互的场景。
内容的提问来源于stack exchange,提问作者lX-Xl
相关产品推荐
相关产品推荐

