Python中如何在starmap多进程间共享矩阵W_hat以并行化parallel_l21函数
解决多进程间共享numpy矩阵W_hat的问题
你的问题核心在于multiprocessing的子进程拥有独立的内存空间,直接传递普通numpy数组时,每个进程都会拿到一份拷贝,修改的只是副本,主进程里的W_hat完全不会被更新。下面给你两种高效的解决方案,适配不同的Python版本:
方法1:使用Python 3.8+的shared_memory(推荐)
shared_memory是Python 3.8引入的标准库,能直接创建跨进程共享的内存块,并且可以无缝转换成numpy数组,效率很高。
修改步骤:
- 先将原W_hat数组存入共享内存,创建一个共享内存对象和对应的numpy数组。
- 在子进程中,通过共享内存的名字重新关联到这个数组。
- 所有进程结束后,手动释放共享内存资源。
修改后的代码:
import numpy as np from multiprocessing import Pool, shared_memory from your_module import rfs # 替换成你的rfs所在模块 def parallel_l21(labels, X, shm_name, shm_shape, shm_dtype, lam, g): # 子进程中关联共享内存 existing_shm = shared_memory.SharedMemory(name=shm_name) W_hat = np.ndarray(shm_shape, dtype=shm_dtype, buffer=existing_shm.buf) g_indxs = np.where(labels == g)[0] # 这里你原代码写的schaefer_labels应该是传入的labels? tmp = rfs(X[g_indxs, 1:].T, X[:, :-1].T, gamma=lam, verbose=False).T tmp[abs(tmp) <= 1e-6] = 0 W_hat[:, g_indxs] = tmp # 关闭共享内存连接(不要在这里unlink,主进程统一处理) existing_shm.close() if __name__ == "__main__": # 假设你的变量已经初始化完成 labels = ... # 你的labels数据 X = ... # 你的X数据 lam = ... # 你的lambda值 W_hat = ... # 初始化好的目标矩阵 # 创建共享内存 shm = shared_memory.SharedMemory(create=True, size=W_hat.nbytes) # 将W_hat的数据复制到共享内存数组 W_shared = np.ndarray(W_hat.shape, dtype=W_hat.dtype, buffer=shm.buf) W_shared[:] = W_hat[:] try: with Pool(processes=12) as pool: # 传递共享内存的名字、形状、数据类型,而非直接传W_hat pool.starmap( parallel_l21, zip( repeat(labels), repeat(X), repeat(shm.name), repeat(W_hat.shape), repeat(W_hat.dtype), repeat(lam), np.unique(labels) ) ) # 将共享内存中的结果同步回主进程的W_hat W_hat[:] = W_shared[:] finally: # 必须释放共享内存,否则会残留系统资源 shm.close() shm.unlink()
方法2:使用multiprocessing.Array(兼容Python 3.7及以下)
如果你的Python版本低于3.8,可以用multiprocessing.Array创建共享内存,再转换成numpy数组:
修改后的代码:
import numpy as np from multiprocessing import Pool, Array from your_module import rfs import ctypes def parallel_l21(labels, X, W_array, W_shape, W_dtype, lam, g): # 将共享Array转换成numpy数组 W_hat = np.frombuffer(W_array.get_obj(), dtype=W_dtype).reshape(W_shape) g_indxs = np.where(labels == g)[0] tmp = rfs(X[g_indxs, 1:].T, X[:, :-1].T, gamma=lam, verbose=False).T tmp[abs(tmp) <= 1e-6] = 0 W_hat[:, g_indxs] = tmp if __name__ == "__main__": labels = ... X = ... lam = ... W_hat = ... # 初始化你的目标矩阵 # 根据numpy dtype匹配对应的ctypes类型 ctype_map = {np.float64: ctypes.c_double, np.float32: ctypes.c_float} c_type = ctype_map[W_hat.dtype] # 创建共享Array,因为各进程修改不同列,不需要锁 W_array = Array(c_type, W_hat.size, lock=False) # 将W_hat的数据复制到共享Array W_shared = np.frombuffer(W_array.get_obj(), dtype=W_hat.dtype).reshape(W_hat.shape) W_shared[:] = W_hat[:] with Pool(processes=12) as pool: pool.starmap( parallel_l21, zip( repeat(labels), repeat(X), repeat(W_array), repeat(W_hat.shape), repeat(W_hat.dtype), repeat(lam), np.unique(labels) ) ) # 同步回主进程的W_hat W_hat[:] = W_shared[:]
关键注意事项:
- 锁的问题:你的代码中每个进程修改的是
W_hat的不同列(g_indxs是每个g对应的独立索引),所以不需要加锁,能大幅提升效率。如果有进程会修改同一区域,一定要加锁(比如把Array的lock参数设为True,或者手动使用Lock)。 - 数据类型匹配:确保共享内存的ctypes类型和numpy的dtype完全一致,否则会出现数据错乱。
- 资源释放:使用
shared_memory时,必须在所有进程结束后调用unlink(),否则系统会残留共享内存资源。
内容的提问来源于stack exchange,提问作者rando
相关产品推荐
相关产品推荐

