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

如何在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

步骤四:启动服务

  1. 启动Redis服务(确保本地Redis已安装并运行)
  2. 启动Celery worker:
celery -A tasks worker --loglevel=info
  1. 启动FastAPI服务:
uvicorn main:app --reload

关键注意事项

  • 资源隔离:如果用GPU训练,需要配置Celery worker的并发数,避免多任务抢占GPU资源;也可以用Celery的路由功能把训练任务分配到特定的GPU节点。
  • 任务持久化:确保Redis的持久化配置开启,避免服务重启丢失任务。
  • 模型版本管理:可以给训练后的模型添加版本号,避免覆盖旧模型导致的回滚问题。
  • 超时设置:在Celery任务中设置超时时间,避免无限期运行的任务占用资源。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.21 07:36:29