使用multiprocessing时DataFrame无法被子进程更新的问题
问题分析与解决
问题原因
在multiprocessing机制中,子进程会获取主进程数据的独立副本,而非直接操作共享内存。你传入子进程的df_1是复制后的对象,子进程内对它的修改仅在自身进程空间生效,主进程的原对象不会受到任何影响。
解决方案
针对你后续要扩展多进程处理数据集的需求,提供几种实用方案:
方案1:用进程池获取返回值(推荐)
通过multiprocessing.Pool管理子进程,异步调用后直接获取处理后的DataFrame,替换主进程的原对象:
from multiprocessing import Pool import time import pandas as pd def test(df): for i, row in df.iterrows(): results = similarity(embeddings, df_acordaos, i, 0) df.at[i, 'SIMILARITY_ALL'] = similarity_formated(results) return df if __name__ == '__main__': start = time.perf_counter() # 后续扩展多进程只需修改processes参数值 with Pool(processes=1) as pool: result = pool.apply_async(test, args=(df_1,)) df_1 = result.get() # 接收子进程处理后的DataFrame finish = time.perf_counter() print(f'Finished in {round(finish-start, 2)} s') print(df_1.columns) # 此时能看到SIMILARITY_ALL列
方案2:用共享内存中转(适合大数据集)
如果数据集体积大,复制副本占用内存过高,可通过multiprocessing.Manager创建可共享的字典对象中转数据:
from multiprocessing import Process, Manager import time import pandas as pd def test(df_shared): # 将共享字典转为DataFrame处理 df = pd.DataFrame(df_shared) for i, row in df.iterrows(): results = similarity(embeddings, df_acordaos, i, 0) df.at[i, 'SIMILARITY_ALL'] = similarity_formated(results) # 处理完转回共享字典 for col in df.columns: df_shared[col] = df[col].tolist() if __name__ == '__main__': start = time.perf_counter() # 把原DataFrame转为可共享字典 df_shared = Manager().dict({col: df_1[col].tolist() for col in df_1.columns}) p = Process(target=test, args=(df_shared,)) p.start() p.join() # 从共享字典恢复为DataFrame df_1 = pd.DataFrame(df_shared) finish = time.perf_counter() print(f'Finished in {round(finish-start, 2)} s') print(df_1.columns)
方案3:用队列传递结果
子进程处理完成后,将DataFrame放入队列,主进程从队列取出结果:
from multiprocessing import Process, Queue import time def test(df, queue): for i, row in df.iterrows(): results = similarity(embeddings, df_acordaos, i, 0) df.at[i, 'SIMILARITY_ALL'] = similarity_formated(results) queue.put(df) if __name__ == '__main__': start = time.perf_counter() queue = Queue() p = Process(target=test, args=(df_1, queue)) p.start() p.join() df_1 = queue.get() # 获取处理后的DataFrame finish = time.perf_counter() print(f'Finished in {round(finish-start, 2)} s') print(df_1.columns)
注意事项
- 所有多进程相关代码必须放在
if __name__ == '__main__':代码块内,避免Windows系统下的进程启动异常。 - 若后续扩展多进程处理多个数据集,方案1的进程池模式最便捷,可通过
map或imap批量提交任务。
内容的提问来源于stack exchange,提问作者Felipe Chermont
相关产品推荐
相关产品推荐

