Python多进程下传递类实例参数的代码加速方法咨询
问题描述
我写了两个多进程处理的代码示例,遇到了明显的性能差异问题:
代码示例1(运行极慢)
from multiprocessing import Pool from functools import partial # 全局定义的类实例,包含多个Pandas DataFrame作为成员变量 metric # 200万元素的列表 variables def get_reward(variable_, metric_, *args): # ... 其他业务逻辑 reward = metric_.some_method(variable_) return reward func_ = partial(get_reward, metric_=metric) mp_workers = Pool(64) map_ = mp_workers.map(func_, variables)
代码示例2(运行速度正常)
from multiprocessing import Pool from functools import partial # 全局定义的类实例,包含多个Pandas DataFrame作为成员变量 metric # 200万元素的列表 variables def get_reward(variable_, *args): # ... 其他业务逻辑 reward = metric.some_method(variable_) return reward mp_workers = Pool(64) map_ = mp_workers.map(get_reward, variables)
两个示例的核心区别:
- 示例2的
get_reward直接使用全局的metric实例 - 示例1通过
partial把metric作为参数传入get_reward
问题:示例1运行速度极慢,推测是因为每个工作进程都会复制庞大的metric实例(包含大体积DataFrame),导致进程初始化和数据传递耗时过长。但实际场景中必须采用示例1的结构——get_reward必须接收类实例作为参数,请问如何优化示例1的性能?
优化方案
1. 使用进程池初始化函数传递全局实例
利用Pool的initializer和initargs参数,在每个工作进程启动时一次性加载metric,避免通过partial重复传递大对象。每个子进程只会初始化一次metric,彻底消除重复复制的开销:
from multiprocessing import Pool from functools import partial # 全局定义的类实例 metric variables # 子进程全局变量,用于存储metric实例 global_metric = None def init_worker(metric_instance): global global_metric global_metric = metric_instance def get_reward(variable_, metric_=None, *args): # 优先使用传入的参数,无参数时用子进程全局实例,满足函数参数要求 use_metric = metric_ if metric_ is not None else global_metric reward = use_metric.some_method(variable_) return reward # 初始化进程池时传入初始化逻辑 mp_workers = Pool(64, initializer=init_worker, initargs=(metric,)) map_ = mp_workers.map(get_reward, variables)
2. 利用Unix系统的fork特性(仅Linux/macOS适用)
Unix系统中multiprocessing.Pool默认用fork模式创建子进程,子进程会直接继承父进程的内存空间(包括全局metric),无需额外复制。可以保留get_reward的参数结构,同时让子进程实际使用继承的全局实例:
from multiprocessing import Pool from functools import partial metric variables def get_reward(variable_, metric_, *args): # 忽略传入的metric_,直接用父进程继承的全局实例,同时保留函数参数结构 reward = metric.some_method(variable_) return reward # 显式指定fork模式(Unix默认),设置maxtasksperchild避免子进程内存泄漏 mp_workers = Pool(64, maxtasksperchild=1000) func_ = partial(get_reward, metric_=metric) map_ = mp_workers.map(func_, variables)
注意:此方法仅适用于只读场景,若父进程修改
metric,子进程不会同步更新。
3. 将DataFrame转换为共享内存对象
如果metric中的DataFrame是只读的,可以将其转为共享内存数组,创建轻量级代理类传递给子进程,避免复制完整数据:
from multiprocessing import Pool, shared_memory from functools import partial import pandas as pd import numpy as np # 将DataFrame转为共享内存对象 def df_to_shared(df): shm = shared_memory.SharedMemory(create=True, size=df.values.nbytes) shared_arr = np.ndarray(df.shape, dtype=df.dtypes, buffer=shm.buf) shared_arr[:] = df.values[:] return shm, shared_arr, df.index, df.columns # 从共享内存恢复DataFrame def shared_to_df(shm_name, shape, dtype, index, columns): existing_shm = shared_memory.SharedMemory(name=shm_name) arr = np.ndarray(shape, dtype=dtype, buffer=existing_shm.buf) return pd.DataFrame(arr, index=index, columns=columns), existing_shm # 创建轻量级Metric代理类,仅存储共享内存信息 class MetricProxy: def __init__(self, shm_name, shape, dtype, index, columns): self.shm_name = shm_name self.shape = shape self.dtype = dtype self.index = index self.columns = columns def some_method(self, variable_): df, shm = shared_to_df(self.shm_name, self.shape, self.dtype, self.index, self.columns) # ... 原some_method业务逻辑 shm.close() # 用完释放共享内存连接 return result # 处理metric中的DataFrame,生成代理实例 shm, shared_arr, idx, cols = df_to_shared(metric.data_df) metric_proxy = MetricProxy(shm.name, shared_arr.shape, shared_arr.dtype, idx, cols) # 传递轻量级代理对象 func_ = partial(get_reward, metric_=metric_proxy) mp_workers = Pool(64) map_ = mp_workers.map(func_, variables) # 所有任务完成后彻底释放共享内存 shm.close() shm.unlink()
4. 使用Manager创建共享实例(谨慎使用)
multiprocessing.Manager可创建跨进程共享对象,但涉及进程间通信,性能不如前三种方法,适合小体积或需要修改的实例:
from multiprocessing import Pool, Manager from functools import partial manager = Manager() # 将metric包装为共享命名空间对象 shared_metric = manager.Namespace() shared_metric.metric = metric variables def get_reward(variable_, metric_, *args): reward = metric_.metric.some_method(variable_) return reward func_ = partial(get_reward, metric_=shared_metric) mp_workers = Pool(64) map_ = mp_workers.map(func_, variables)
内容的提问来源于stack exchange,提问作者Yanghoon
相关产品推荐
相关产品推荐

