如何确保Python Multiprocessing Pool中所有Worker至少执行一次任务?
解决方案:确保Multiprocessing Pool中每个Worker执行Commit并充分利用所有进程
问题核心分析
- 任务分配机制限制:Python
multiprocessing.Pool的调度策略会优先复用空闲Worker,而非强制平均分配任务,导致少量Worker即可处理完所有任务,其余Worker长期空闲,无法触发Commit操作。 - 测试代码的隐性问题:使用
imap时未遍历返回的迭代器,任务未被完全触发执行,这也是仅少数Worker运行的原因之一。
解决方法
方法1:用apply_async强制每个Worker执行Commit
通过提交与Worker数量相等的Commit任务,确保每个Worker至少被分配一次任务,从而触发所有数据库连接的Commit操作。
import os import multiprocessing import psycopg2 import traceback import time def conn_db(): try: conn = psycopg2.connect(database="dbname", user="username", password="pass", host="127.0.0.1", port="5432") return conn except: print(traceback.format_exc()) def init_proc(): global conn conn = conn_db() global cursor cursor = conn.cursor() def update(record): # 示例DML操作 cursor.execute("UPDATE table SET col = %s WHERE id = %s", (record[1], record[0])) def func_commit(_): global conn conn.commit() print(f"Process {os.getpid()} committed transaction") def exec_parallel_update(records): try: pool = multiprocessing.Pool(8, initializer=init_proc) t1_b = time.time() # 执行所有更新任务,必须遍历结果触发执行 results = pool.map(update, records) # 提交与Worker数量相等的Commit任务 commit_tasks = [] for _ in range(8): task = pool.apply_async(func_commit, (None,)) commit_tasks.append(task) # 等待所有Commit任务完成 for task in commit_tasks: task.get() pool.close() pool.join() t1_runtime = time.time() - t1_b print(f'Updated {len(records)} records, runtime: {t1_runtime:.2f}s') except: print(traceback.format_exc())
方法2:手动创建Process,完全控制任务分配
放弃Pool,手动创建与Worker数量一致的进程,每个进程负责处理一部分数据并自行Commit,从根源上确保每个进程都执行Commit。
import multiprocessing import psycopg2 import traceback import time def conn_db(): try: conn = psycopg2.connect(database="dbname", user="username", password="pass", host="127.0.0.1", port="5432") return conn except: print(traceback.format_exc()) def process_chunk(chunk): conn = conn_db() cursor = conn.cursor() for record in chunk: cursor.execute("UPDATE table SET col = %s WHERE id = %s", (record[1], record[0])) conn.commit() conn.close() print(f"Process {multiprocessing.current_process().pid} finished chunk and committed") def exec_parallel_update(records): try: t1_b = time.time() # 将数据拆分为8个分片 chunk_size = len(records) // 8 chunks = [records[i*chunk_size : (i+1)*chunk_size] for i in range(8)] # 处理剩余数据 if len(records) % 8 != 0: chunks[-1].extend(records[8*chunk_size:]) # 创建并启动进程 processes = [] for chunk in chunks: p = multiprocessing.Process(target=process_chunk, args=(chunk,)) processes.append(p) p.start() # 等待所有进程完成 for p in processes: p.join() t1_runtime = time.time() - t1_b print(f'Updated {len(records)} records, runtime: {t1_runtime:.2f}s') except: print(traceback.format_exc())
方法3:利用进程退出钩子自动Commit
在Worker初始化时注册退出钩子,当Worker进程退出时自动执行Commit操作,适合无需手动控制Commit时机的场景。
import multiprocessing import psycopg2 import traceback import time import atexit def conn_db(): try: conn = psycopg2.connect(database="dbname", user="username", password="pass", host="127.0.0.1", port="5432") return conn except: print(traceback.format_exc()) def init_proc(): global conn conn = conn_db() global cursor cursor = conn.cursor() # 注册进程退出钩子,退出时自动Commit atexit.register(lambda: conn.commit()) def update(record): cursor.execute("UPDATE table SET col = %s WHERE id = %s", (record[1], record[0])) def exec_parallel_update(records): try: pool = multiprocessing.Pool(8, initializer=init_proc) t1_b = time.time() pool.map(update, records) pool.close() pool.join() # Worker进程在此后退出,触发Commit t1_runtime = time.time() - t1_b print(f'Updated {len(records)} records, runtime: {t1_runtime:.2f}s') except: print(traceback.format_exc())
修复测试代码的Worker利用率问题
你的测试代码中imap返回的是迭代器,必须遍历才能触发任务执行,修改后即可让所有Worker参与:
import os import multiprocessing import time def init_proc(): global conn # 模拟数据库连接初始化 conn = None def double2(i): print(f"I'm process:{os.getpid()}, {multiprocessing.current_process()}") return i*2 def exec_ten_parallel(num_parallel, records): try: pool = multiprocessing.Pool(8, initializer=init_proc) t1_b = time.time() # 必须遍历imap结果触发任务执行 results = pool.imap(double2, range(16), 1) for res in results: pass pool.close() pool.join() t1_runtime = time.time() - t1_b except: import traceback print(traceback.format_exc()) exec_ten_parallel(8, [])
修改后会看到所有8个Worker都被调用执行任务。
内容的提问来源于stack exchange,提问作者Jerry Chen
相关产品推荐
相关产品推荐

