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

基于Flask+SocketIO的ML训练服务多线程/进程问题咨询

解决方案:主进程+训练子进程+IPC队列架构

针对你的问题,最优方案是采用主进程负责IO通信、子进程执行训练、进程间队列传递进度的架构,既解决线程无法强制终止的问题,又满足仅主进程更新客户端的要求,无需额外在进程内创建冗余线程。

核心逻辑

  1. 主进程:处理HTTP请求、SocketIO客户端通信、管理训练子进程,通过监听进程间队列获取训练进度并推送给客户端。
  2. 训练子进程:独立执行耗时的model.train(...)任务,仅负责将训练进度写入进程间队列,不直接操作SocketIO。
  3. 进程间通信(IPC):用multiprocessing.Queue作为进度传递的桥梁,子进程写入进度,主进程读取后推送。

具体实现步骤

1. 初始化核心组件

  • 主进程创建SocketIO实例、训练任务注册表(存储任务ID与对应子进程的映射)、进程间进度队列。
  • 主进程启动一个后台线程,持续监听进度队列,一旦有新进度就通过SocketIO推送给对应客户端。

2. 触发训练流程

  • 客户端上传文件后,主进程生成唯一任务ID,创建训练子进程,传入任务ID、数据路径、进度队列。
  • 子进程启动后,将其存入训练任务注册表,同时返回任务ID给客户端,用于后续进度接收和终止请求。

3. 进度推送机制

  • 子进程在训练循环(或model.train()的回调钩子)中,定期将epoch、loss、当前步骤等进度数据写入队列。
  • 主进程的监听线程读取队列中的进度数据,通过SocketIO的房间机制(每个任务对应一个房间)推送给已加入该房间的客户端。

4. 强制终止训练

  • 客户端发送停止请求时,主进程根据任务ID找到对应子进程,调用terminate()强制终止(多进程原生支持),清理注册表后通知客户端训练已停止。

代码示例

主进程代码

from flask import Flask, request
from flask_socketio import SocketIO, emit, join_room
import multiprocessing as mp
import time
import os

app = Flask(__name__)
socketio = SocketIO(app, cors_allowed_origins="*")

# 存储训练任务:key=任务ID,value=子进程实例
training_tasks = {}
# 进程间进度传递队列
progress_queue = mp.Queue()

def progress_listener():
    """主进程后台线程:监听进度队列并推送至客户端"""
    while True:
        if not progress_queue.empty():
            task_id, progress_data = progress_queue.get()
            # 向对应任务房间推送进度
            socketio.emit('training_progress', {
                'task_id': task_id,
                'data': progress_data
            }, room=task_id)
        time.sleep(0.1)

# 启动进度监听线程
socketio.start_background_task(target=progress_listener)

def train_subprocess(task_id, data_path, queue):
    """子进程:执行训练任务并写入进度"""
    # 模拟数据加载与模型初始化(替换为你的实际代码)
    print(f"子进程启动训练任务:{task_id}")
    model = type('MockModel', (), {})()
    setattr(model, 'train_step', lambda x: 1.0 - (x/10))

    # 模拟训练循环(替换为你的model.train()逻辑,需插入进度回调)
    for epoch in range(10):
        # 检查进程是否已被终止,避免无效执行
        if mp.current_process().exitcode is not None:
            break
        # 执行单轮训练
        current_loss = model.train_step(epoch)
        # 写入进度数据
        queue.put((task_id, {
            'epoch': epoch + 1,
            'loss': round(current_loss, 4),
            'status': 'running'
        }))
        time.sleep(1)  # 模拟训练耗时

    # 训练结束/终止后发送状态
    final_status = 'completed' if mp.current_process().exitcode is None else 'stopped'
    queue.put((task_id, {
        'status': final_status,
        'epoch': epoch + 1 if final_status == 'completed' else epoch
    }))

@app.route('/start-training', methods=['POST'])
def start_training():
    # 接收上传文件
    if 'data_file' not in request.files:
        return {'error': 'No file uploaded'}, 400
    file = request.files['data_file']
    temp_dir = './temp_data'
    os.makedirs(temp_dir, exist_ok=True)
    data_path = os.path.join(temp_dir, file.filename)
    file.save(data_path)

    # 生成唯一任务ID
    task_id = f"train_task_{int(time.time())}"
    # 创建并启动训练子进程
    train_process = mp.Process(
        target=train_subprocess,
        args=(task_id, data_path, progress_queue)
    )
    train_process.start()
    training_tasks[task_id] = train_process

    return {'task_id': task_id}, 200

@app.route('/stop-training/<task_id>', methods=['POST'])
def stop_training(task_id):
    if task_id not in training_tasks:
        return {'error': 'Task not found'}, 404
    process = training_tasks[task_id]
    if process.is_alive():
        process.terminate()
        process.join()
    del training_tasks[task_id]
    # 通知客户端终止状态
    socketio.emit('training_status', {
        'task_id': task_id,
        'status': 'stopped'
    }, room=task_id)
    return {'status': 'success'}, 200

@socketio.on('join-task-room')
def handle_join_room(task_id):
    """客户端加入对应任务的SocketIO房间"""
    join_room(task_id)
    emit('room-joined', {'msg': f'Successfully joined task {task_id}'})

if __name__ == '__main__':
    socketio.run(app, host='0.0.0.0', port=5000, debug=True)

关键优势

  • 强制终止:子进程可通过terminate()直接终止,解决线程无法强制停止的痛点。
  • 解耦设计:训练进程专注于训练逻辑,主进程专注于客户端通信,职责清晰。
  • 可靠性:进程间队列是线程/进程安全的,不会出现进度丢失或混乱的情况。

内容的提问来源于stack exchange,提问作者vct12345

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.23 09:35:03