Python多进程/多线程中异常与CTRL+C处理及相关技术问题
多进程/线程训练测试框架问题解决方案
核心问题:优雅处理CTRL+C终止
当前程序无法捕获CTRL+C是因为主线程被join()阻塞,且子进程的信号处理干扰了主线程。解决方法:
- 在
__main__中捕获KeyboardInterrupt,通过线程安全的事件触发全局终止 - 使用
threading.Event替代原terminate变量,避免跨线程竞态条件 - 主线程捕获中断后,主动通知所有子线程和进程退出
额外问题解答
1. 训练线程结束时的优雅终止方式
用threading.Event替代nonlocal terminate变量是更优方案:
Event是线程安全的,无需手动处理竞态- 训练线程结束时调用
terminate_event.set(),所有监听的线程/进程能立即感知并启动退出流程 - 避免了原代码中依赖共享变量的潜在问题
2. 单个进程池 vs 两个进程池
| 方案 | 优势 | 适用场景 |
|---|---|---|
| 单个进程池 | 实现简单,进程复用率高 | 采集和测试任务资源需求相似(均为CPU密集型) |
| 两个进程池 | 资源隔离,可分别控制两类任务的进程数,避免互相抢占 | 采集和测试资源需求差异大(如采集用CPU、测试用GPU),或需要单独启停某一类任务 |
建议优先用两个进程池,灵活性更高,后续扩展更方便。
3. 让对应线程捕获子进程异常
利用multiprocessing.Pool.apply_async的error_callback参数,将子进程异常定向传递给对应线程:
- 采集进程的异常通过
error_callback通知训练线程,触发全局终止 - 测试进程的异常通过
error_callback仅通知测试线程,不影响训练流程
修改后的完整代码
import threading import logging import traceback import torch import time from torch import multiprocessing as mp try: mp.set_start_method('spawn') except RuntimeError: pass shandle = logging.StreamHandler() log = logging.getLogger('rl') log.propagate = False log.addHandler(shandle) log.setLevel(logging.INFO) def collect(id, queue, data_collect): log.info('Collect %i started ...', id) try: while True: idx = queue.get() if idx is None: break data_collect[idx] = torch.rand(1) queue.task_done() # 实际业务逻辑 except Exception as e: log.error('Exception in collect process %i', id) traceback.print_exc() raise e def test(id, queue, data_test): log.info('Test %i started ...', id) try: while True: idx = queue.get() if idx is None: break data_test[idx] = torch.rand(1) queue.task_done() # 实际业务逻辑 except Exception as e: log.error('Exception in test process %i', id) traceback.print_exc() raise e def run(terminate_event): steps = 0 num_collect_procs = 3 num_test_procs = 2 max_steps = 10 data_collect = torch.zeros(num_collect_procs).share_memory_() data_test = torch.zeros(num_test_procs).share_memory_() manager = mp.Manager() # 两个进程池分别管理采集、测试任务 collect_pool = mp.Pool(num_collect_procs) test_pool = mp.Pool(num_test_procs) collect_queue = manager.JoinableQueue() test_queue = manager.JoinableQueue() # 训练线程异常标志 train_error = threading.Event() # 采集进程异常回调:触发训练终止+全局终止 def collect_error_callback(e): log.error('Collect process error, triggering train termination') train_error.set() terminate_event.set() # 测试进程异常回调:仅终止测试流程 def test_error_callback(e): log.error('Test process error, stopping testing') test_error.set() # 启动采集进程 for i in range(num_collect_procs): collect_pool.apply_async(collect, args=(i, collect_queue, data_collect), error_callback=collect_error_callback) # 启动测试进程 test_error = threading.Event() for i in range(num_test_procs): test_pool.apply_async(test, args=(i, test_queue, data_test), error_callback=test_error_callback) # 训练线程逻辑 def run_train(): nonlocal steps log.info('Training thread started ...') while steps < max_steps and not terminate_event.is_set() and not train_error.is_set(): try: for idx in range(num_collect_procs): collect_queue.put(idx) collect_queue.join() time.sleep(0.1) log.info('Training, %i %f', steps, data_collect.sum()) steps += 1 except Exception as e: log.error('Training thread exception') traceback.print_exc() train_error.set() terminate_event.set() break # 通知采集进程退出 for _ in range(num_collect_procs): collect_queue.put(None) collect_queue.join() log.info('Training done') # 测试线程逻辑 def run_test(): nonlocal steps log.info('Testing thread started ...') while steps < max_steps and not terminate_event.is_set() and not test_error.is_set(): try: for idx in range(num_test_procs): test_queue.put(idx) test_queue.join() time.sleep(0.1) log.info('Testing, %i %f', steps, data_test.sum()) except Exception as e: log.error('Testing thread exception') traceback.print_exc() test_error.set() break # 通知测试进程退出 for _ in range(num_test_procs): test_queue.put(None) test_queue.join() log.info('Testing done') learning_thread = threading.Thread(target=run_train, name='train') learning_thread.start() testing_thread = threading.Thread(target=run_test, name='test') testing_thread.start() # 训练线程结束则触发全局终止 learning_thread.join() terminate_event.set() testing_thread.join() # 关闭进程池 collect_pool.terminate() collect_pool.join() test_pool.terminate() test_pool.join() if __name__ == '__main__': terminate_event = threading.Event() try: run(terminate_event) except KeyboardInterrupt: log.info('Received CTRL+C, terminating all processes...') terminate_event.set()
关键改进点
- 用
threading.Event实现线程安全的终止信号传递 - 两个进程池实现采集、测试任务的资源隔离
- 通过
error_callback定向传递子进程异常,满足业务需求 - 捕获
KeyboardInterrupt后触发优雅退出流程
内容的提问来源于stack exchange,提问作者Simon
相关产品推荐
相关产品推荐

