Python multiprocessing如何每n次全局迭代在独立线程执行函数
问题背景
已搭建多进程工作任务,初始代码如下:
global_iteration = mp.Value('i', 0) workers = [Worker(global_iteration) for _ in range(num_threads)] for w in workers: w.daemon=True w.start() [w.join() for w in workers]
每个Worker实例执行完业务操作后,会将global_iteration全局迭代计数加1,原Worker类定义如下(存在几处语法笔误,后续实现中已修正):
import multiprocessing as mp class Worker(mp.Process): def __init__(global_iteration): super(Worker, self).__init__() self.global_iteration = global_iteration def update_global_iteration(): with self.global_iteration.get_lock(): self.global_iteration+=1 def run(): ### Do something here ### self.update_global_iteration()
具体问题
需要实现逻辑:每完成n次global_iteration全局迭代,就在独立线程中运行一次指定函数。示例待执行函数如下:
def print_global_iterations(global_iterations): print('Workers are currently on global iteration {}'.format(global_iterations))
实现方案
核心思路是在主进程启动单独的守护线程做计数监听,避免在多个Worker子进程中分散判断导致回调重复触发:
- 监听线程维护上一次触发回调的计数阈值,循环读取跨进程共享的迭代计数值
- 当计数达到下一个n的整数倍阈值时,启动独立线程执行目标回调,更新下一次触发阈值
- 所有Worker进程执行完毕后,监听线程自动退出
完整可运行代码:
import multiprocessing as mp import threading import time def print_global_iterations(global_iterations): print('Workers are currently on global iteration {}'.format(global_iterations)) class Worker(mp.Process): def __init__(self, global_iteration): super().__init__() self.global_iteration = global_iteration def update_global_iteration(self): with self.global_iteration.get_lock(): self.global_iteration.value += 1 def run(self): # 替换为实际业务逻辑 time.sleep(0.1) self.update_global_iteration() def iteration_monitor(shared_count, trigger_step, callback): next_trigger_val = trigger_step total_task_num = num_workers while True: current = shared_count.value if current >= next_trigger_val: # 启动独立线程执行回调,不阻塞监听逻辑 threading.Thread(target=callback, args=(current,), daemon=True).start() next_trigger_val += trigger_step # 短休眠降低CPU占用,可根据实时性要求调整间隔 time.sleep(0.01) # 所有任务完成后退出监听 if current >= total_task_num: break if __name__ == "__main__": num_workers = 20 # 每完成5次迭代触发一次回调,可自行修改n值 trigger_n = 5 global_iteration = mp.Value('i', 0) # 启动监听线程 monitor_thread = threading.Thread( target=iteration_monitor, args=(global_iteration, trigger_n, print_global_iterations), daemon=True ) monitor_thread.start() # 启动所有工作进程 workers = [Worker(global_iteration) for _ in range(num_workers)] for w in workers: w.start() for w in workers: w.join() monitor_thread.join()
关键注意事项
- 操作
mp.Value类型的跨进程共享变量时,必须通过.value属性读写值,直接对变量对象做赋值/加减操作不会修改实际存储的共享值 - 计数累加操作必须放在锁保护的代码块内执行,避免多进程并发写入导致计数不准确
- 不要在Worker子进程内实现触发判断,多进程同时读取到阈值时会重复执行回调函数
- 监听线程增加短休眠避免CPU空转,对触发实时性要求高可以把休眠间隔调小
内容的提问来源于stack exchange,提问作者SRS
相关产品推荐
相关产品推荐

