Python多进程中共享大DataFrame参数的内存优化问题
解决大DataFrame多进程重复传参的内存浪费问题
方案一:利用进程池初始化函数传递一次DataFrame
通过multiprocessing.Pool的initializer和initargs参数,在每个子进程启动时仅加载一次DataFrame作为子进程的全局变量,避免每次任务都传递完整的df副本。
修改后的代码:
import multiprocessing from multiprocessing import Pool from tqdm import tqdm # 定义子进程全局变量,用于存储DataFrame global_df = None def init_worker(df): """子进程初始化函数,将传入的df赋值给全局变量""" global global_df global_df = df def function_to_run(lvl, timescale, adm_code): """修改后的业务函数,直接使用子进程的全局global_df""" # 这里写你的业务逻辑,例如: subset = global_df[global_df[lvl] == adm_code] # ...其他处理代码 def parallel_func(args): lvl, timescale, adm_code = args return function_to_run(lvl, timescale, adm_code) def main(): for country in all_countries: # 仅执行一次SQL查询获取df df = SqlManager().sql_query('MyDb', f"SELECT * FROM MyTable WHERE Country='{country}'") # 构造参数列表,不再包含df args = [ (lvl, timescale, adm_code) for lvl in ['Country', 'Region', 'County'] for timescale in ['month', 'week'] for adm_code in list(df[lvl].unique()) ] # 创建进程池时指定初始化函数,每个子进程启动时加载df with Pool(processes=multiprocessing.cpu_count(), initializer=init_worker, initargs=(df,)) as pool: entry_list = list(tqdm(pool.imap(parallel_func, args), total=len(args)))
核心优势:
- 每个子进程仅会拷贝一次df到自身内存空间,假设CPU核心数为8,仅产生8份df副本,相比原方案的上万次拷贝,内存占用大幅降低。
- 无需在子进程中重复执行SQL查询,避免了数据库连接和查询的性能损耗。
方案二:使用共享内存(适合超大规模DataFrame)
如果df规模达到数亿行,方案一的多份副本仍可能占用过多内存,可以使用Pandas的共享内存特性(需Pandas 1.3.0+版本),让所有子进程共享同一块内存中的DataFrame,完全避免内存拷贝。
修改后的代码:
import multiprocessing from multiprocessing import Pool from tqdm import tqdm import pandas as pd def function_to_run(args): lvl, timescale, adm_code, shm_info = args # 从共享内存中恢复DataFrame df = pd.read_shared_memory(shm_info) # 业务逻辑处理 subset = df[df[lvl] == adm_code] # ...其他处理代码 def parallel_func(args): return function_to_run(args) def main(): for country in all_countries: df = SqlManager().sql_query('MyDb', f"SELECT * FROM MyTable WHERE Country='{country}'") # 将df写入共享内存,返回共享内存标识信息 shm_info = df.to_shared_memory() # 构造参数列表,传递共享内存信息而非完整df args = [ (lvl, timescale, adm_code, shm_info) for lvl in ['Country', 'Region', 'County'] for timescale in ['month', 'week'] for adm_code in list(df[lvl].unique()) ] with Pool(processes=multiprocessing.cpu_count()) as pool: entry_list = list(tqdm(pool.imap(parallel_func, args), total=len(args))) # 手动释放共享内存,避免内存泄漏 shm_info.close() shm_info.unlink()
注意事项:
- 共享内存中的DataFrame为只读状态,若子进程需要修改数据,需先创建本地副本(你的场景中仅需读取,无需修改)。
- 必须在任务执行完毕后调用
close()和unlink()释放共享内存,否则会造成内存泄漏。
原方案内存浪费的根源:
原代码中,args列表的每个元组都包含完整的df对象,当args包含上万条数据时,相当于在内存中存储了上万份df副本。同时Python多进程传递参数时会通过pickle序列化,每个任务都要序列化一次df,既消耗内存又占用CPU资源。
内容的提问来源于stack exchange,提问作者kilag
相关产品推荐
相关产品推荐

