You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

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数组,效率很高。

修改步骤:

  1. 先将原W_hat数组存入共享内存,创建一个共享内存对象和对应的numpy数组。
  2. 在子进程中,通过共享内存的名字重新关联到这个数组。
  3. 所有进程结束后,手动释放共享内存资源。

修改后的代码:

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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.04.29 19:12:35