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

Python多进程实现共享数据矩阵的最小二乘交叉验证优化

Python多进程共享内存实现索引式交叉验证

原生multiprocessing完全支持你的需求,下面提供两种高效实现方案,解决数据重复传递导致的内存占用问题:

方案1:利用进程池初始化器传递共享数据(适合中小规模数据)

通过Pool的initializer参数,在子进程启动时将X和Y注入为全局变量,子进程仅需接收索引子集即可计算,避免重复传递完整数据。

import numpy as np
from multiprocessing import Pool

# 子进程全局共享变量,存储X和Y
_shared_X = None
_shared_Y = None

def init_worker(X, Y):
    """初始化工作进程,绑定共享数据到全局变量"""
    global _shared_X, _shared_Y
    _shared_X = X
    _shared_Y = Y

def solve_lsqr_with_indices(indices):
    """接收索引子集,使用共享数据求解最小二乘"""
    X_subset = _shared_X[indices]
    Y_subset = _shared_Y[indices]
    # 替换为你的最小二乘求解逻辑
    w, _, _, _ = np.linalg.lstsq(X_subset, Y_subset, rcond=None)
    return w

if __name__ == "__main__":
    # 示例训练数据
    X = np.random.rand(1000, 10)
    Y = np.random.rand(1000, 1)

    # 生成自定义滞后采样的索引集(替换为你的实际逻辑)
    cv_indices = [
        np.arange(0, 300),
        np.arange(200, 500),
        np.arange(400, 700)
    ]

    # 创建进程池并注入共享数据
    with Pool(processes=4, initializer=init_worker, initargs=(X, Y)) as pool:
        results = pool.map(solve_lsqr_with_indices, cv_indices)

    # 处理交叉验证结果
    for idx, w in enumerate(results):
        print(f"第{idx+1}份子集的参数w:\n{w}")

说明:Unix系统下基于fork机制,子进程会继承父进程内存(写时复制,仅修改时才复制数据);Windows下基于spawn机制,数据会通过初始化器传递给每个子进程,适合中小规模数据场景。

方案2:使用共享内存块(适合大规模数据)

通过multiprocessing.shared_memory创建真正的内存共享块,所有子进程直接访问同一份内存,彻底避免数据复制,大幅降低内存占用。

import numpy as np
from multiprocessing import Pool, shared_memory

def solve_lsqr_with_shmem(indices, x_shm_name, x_shape, x_dtype, y_shm_name, y_shape, y_dtype):
    """通过共享内存名称获取数据,处理索引子集"""
    # 连接到共享内存块
    x_shm = shared_memory.SharedMemory(name=x_shm_name)
    X = np.ndarray(x_shape, dtype=x_dtype, buffer=x_shm.buf)
    
    y_shm = shared_memory.SharedMemory(name=y_shm_name)
    Y = np.ndarray(y_shape, dtype=y_dtype, buffer=y_shm.buf)

    # 求解最小二乘
    X_subset = X[indices]
    Y_subset = Y[indices]
    w, _, _, _ = np.linalg.lstsq(X_subset, Y_subset, rcond=None)

    # 关闭共享内存连接(主进程负责清理)
    x_shm.close()
    y_shm.close()
    return w

if __name__ == "__main__":
    # 示例大规模训练数据
    X = np.random.rand(10000, 100)
    Y = np.random.rand(10000, 1)

    # 创建X的共享内存块并写入数据
    x_shm = shared_memory.SharedMemory(create=True, size=X.nbytes)
    X_shared = np.ndarray(X.shape, dtype=X.dtype, buffer=x_shm.buf)
    X_shared[:] = X[:]

    # 创建Y的共享内存块并写入数据
    y_shm = shared_memory.SharedMemory(create=True, size=Y.nbytes)
    Y_shared = np.ndarray(Y.shape, dtype=Y.dtype, buffer=y_shm.buf)
    Y_shared[:] = Y[:]

    # 生成自定义滞后采样的索引集
    cv_indices = [
        np.arange(0, 3000),
        np.arange(2000, 5000),
        np.arange(4000, 7000)
    ]

    # 整理传递给子进程的共享内存参数
    shm_task_args = [
        (indices, x_shm.name, X.shape, X.dtype, y_shm.name, Y.shape, Y.dtype)
        for indices in cv_indices
    ]

    with Pool(processes=4) as pool:
        results = pool.starmap(solve_lsqr_with_shmem, shm_task_args)

    # 清理共享内存,避免内存泄漏
    x_shm.close()
    x_shm.unlink()
    y_shm.close()
    y_shm.unlink()

    # 处理结果
    for idx, w in enumerate(results):
        print(f"第{idx+1}份子集的参数w:\n{w}")

说明:此方案中X和Y仅在内存中存储一份,所有子进程直接读写同一块内存,适合超大规模数据场景。注意必须由主进程执行unlink操作清理共享内存。

关键注意事项

  • 自定义滞后采样的索引集需根据你的业务逻辑生成,确保符合交叉验证的子采样规则。
  • 每个子进程的内部变量相互独立,仅X和Y为共享数据,不会出现状态混乱问题。

内容的提问来源于stack exchange,提问作者arash

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.31 10:15:35