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

如何确保Python Multiprocessing Pool中所有Worker至少执行一次任务?

解决方案:确保Multiprocessing Pool中每个Worker执行Commit并充分利用所有进程

问题核心分析

  1. 任务分配机制限制:Python multiprocessing.Pool 的调度策略会优先复用空闲Worker,而非强制平均分配任务,导致少量Worker即可处理完所有任务,其余Worker长期空闲,无法触发Commit操作。
  2. 测试代码的隐性问题:使用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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.18 21:30:51