基于Flask+SocketIO的ML训练服务多线程/进程问题咨询
解决方案:主进程+训练子进程+IPC队列架构
针对你的问题,最优方案是采用主进程负责IO通信、子进程执行训练、进程间队列传递进度的架构,既解决线程无法强制终止的问题,又满足仅主进程更新客户端的要求,无需额外在进程内创建冗余线程。
核心逻辑
- 主进程:处理HTTP请求、SocketIO客户端通信、管理训练子进程,通过监听进程间队列获取训练进度并推送给客户端。
- 训练子进程:独立执行耗时的
model.train(...)任务,仅负责将训练进度写入进程间队列,不直接操作SocketIO。 - 进程间通信(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
相关产品推荐
相关产品推荐

