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

Python多进程/多线程中异常与CTRL+C处理及相关技术问题

多进程/线程训练测试框架问题解决方案

核心问题:优雅处理CTRL+C终止

当前程序无法捕获CTRL+C是因为主线程被join()阻塞,且子进程的信号处理干扰了主线程。解决方法:

  1. 在__main__中捕获KeyboardInterrupt,通过线程安全的事件触发全局终止
  2. 使用threading.Event替代原terminate变量,避免跨线程竞态条件
  3. 主线程捕获中断后,主动通知所有子线程和进程退出

额外问题解答

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.23 20:27:49