如何在FastAPI端点中集成耗时的深度学习模型训练任务?
用FastAPI实现耗时深度学习模型的异步重训练端点
网上很多资源讲如何用FastAPI基于预训练静态模型做预测,但如果要创建端点启动基于新数据的模型重训练,FastAPI的BackgroundTasks并不适合这类耗时数分钟、占用大量CPU/GPU的深度学习训练任务。下面给出具体的扩展方案,基于原有的代码框架进行改造。
核心思路
BackgroundTasks的任务和请求上下文绑定,服务重启或请求超时会导致任务中断,且占用主线程资源影响其他请求。因此我们需要用独立的异步任务队列来解耦训练任务和API服务,推荐使用Celery搭配Redis(或RabbitMQ)作为消息代理和结果存储,实现训练任务的异步执行、状态跟踪和结果持久化。
步骤一:安装依赖
需要安装以下包:
pip install fastapi uvicorn celery redis numpy scikit-learn # 根据你的深度学习框架替换scikit-learn,比如tensorflow/pytorch
步骤二:配置Celery任务队列
创建tasks.py文件,定义训练任务和Celery配置:
import pickle import numpy as np from celery import Celery from sklearn.ensemble import RandomForestClassifier # 替换为你的深度学习模型类 # 初始化Celery,用Redis作为消息代理和结果存储 celery_app = Celery( "training_tasks", broker="redis://localhost:6379/0", backend="redis://localhost:6379/0" ) # 模拟加载新训练数据(实际场景可以从数据库/存储读取) def load_new_training_data(): # 示例数据,替换为你的真实数据加载逻辑 X = np.random.rand(1000, 3) y = np.random.randint(0, 2, size=1000) return X, y @celery_app.task(bind=True) def retrain_model_task(self): try: # 加载新训练数据 X_train, y_train = load_new_training_data() # 初始化或加载现有模型(如果需要基于旧模型微调) model = RandomForestClassifier() # 替换为你的深度学习模型初始化逻辑 # 如果要微调旧模型,可以从文件加载后继续训练 # with open("../app/model.pkl", "rb") as f: # model = pickle.load(f) # 开始训练 self.update_state(state='PROGRESS', meta={'status': '训练中...'}) model.fit(X_train, y_train) # 替换为你的深度学习训练代码 # 保存新模型 with open("../app/model.pkl", "wb") as f: pickle.dump(model, f) return {'status': '训练完成', 'model_path': '../app/model.pkl'} except Exception as e: return {'status': '训练失败', 'error': str(e)}
步骤三:改造FastAPI主程序
修改main.py,添加训练启动、状态查询端点,并处理模型加载的并发问题:
import numpy as np import pickle from fastapi import FastAPI, HTTPException from pydantic import BaseModel import os from threading import Lock from tasks import retrain_model_task, celery_app app = FastAPI(title="模型预测与重训练服务") # 模型锁,避免训练替换模型时和预测请求并发冲突 model_lock = Lock() model = None class Datapoint(BaseModel): feature1: float feature2: float feature3: float @app.on_event("startup") def load_model(): global model model_path = "../app/model.pkl" if os.path.exists(model_path): with open(model_path, "rb") as file: model = pickle.load(file) else: # 如果没有预训练模型,初始化一个默认模型 from sklearn.ensemble import RandomForestClassifier model = RandomForestClassifier() with open(model_path, "wb") as f: pickle.dump(model, f) @app.get("/") async def root(): return {"message": "模型预测与重训练服务运行中"} @app.post("/predict") def predict(data: Datapoint): if model is None: raise HTTPException(status_code=500, detail="模型未加载") with model_lock: data_point = np.array([[data.feature1, data.feature2, data.feature3]]) pred = model.predict(data_point).tolist()[0] return {"Prediction": pred} @app.post("/retrain") def start_retrain(): # 提交训练任务到Celery队列 task = retrain_model_task.delay() return {"task_id": task.id, "status": "训练任务已提交"} @app.get("/retrain/status/{task_id}") def get_retrain_status(task_id: str): task = retrain_model_task.AsyncResult(task_id) if task.state == 'PENDING': response = {"status": "任务等待执行"} elif task.state == 'PROGRESS': response = {"status": task.info['status']} elif task.state == 'SUCCESS': response = {"status": task.info['status'], "model_path": task.info['model_path']} # 可选:自动重新加载新模型 with model_lock: global model with open("../app/model.pkl", "rb") as f: model = pickle.load(f) else: response = {"status": "任务失败", "error": task.info.get('error', '未知错误')} return response
步骤四:启动服务
- 启动Redis服务(确保本地Redis已安装并运行)
- 启动Celery worker:
celery -A tasks worker --loglevel=info
- 启动FastAPI服务:
uvicorn main:app --reload
关键注意事项
- 资源隔离:如果用GPU训练,需要配置Celery worker的并发数,避免多任务抢占GPU资源;也可以用Celery的路由功能把训练任务分配到特定的GPU节点。
- 任务持久化:确保Redis的持久化配置开启,避免服务重启丢失任务。
- 模型版本管理:可以给训练后的模型添加版本号,避免覆盖旧模型导致的回滚问题。
- 超时设置:在Celery任务中设置超时时间,避免无限期运行的任务占用资源。
内容的提问来源于stack exchange,提问作者Patrick
相关产品推荐
相关产品推荐

