FastAPI 应在启动时加载模型而非每次请求:用 lifespan 机制全局加载,避免 on_event;输入需用 Pydantic 校验并转为 numpy 数组;输出须显式转为 Python 原生类型。

FastAPI 启动时加载模型比每次请求都加载快得多
直接在 predict 路由函数里用 joblib.load() 加载模型,会导致每次请求都反序列化一次,CPU 和 I/O 开销叠加,QPS 瞬间掉一半。正确做法是应用启动时一次性加载,并存为全局变量或依赖项。
推荐用 FastAPI 的 lifespan 机制(Python 3.8+)管理生命周期:
from fastapi import FastAPI
from sklearn.ensemble import RandomForestClassifier
import joblib
<p>model = None</p><p>async def lifespan(app: FastAPI):
global model
model = joblib.load("model.pkl") # 启动时加载一次
yield
model = None # 可选:退出时清理</p><p>app = FastAPI(lifespan=lifespan)
</p>
- 避免用
on_event("startup")—— 已被lifespan替代,旧写法在最新 FastAPI 中会警告 - 如果模型大于 200MB,考虑用
mmap_mode="r"参数传给joblib.load(),减少内存拷贝 - 别把模型放在
Depends()函数里返回——那只是“每次请求都调一次”,没解决根本问题
输入数据必须和训练时的 shape/dtype 完全一致
FastAPI 默认把 JSON 解析成 dict 或 list,而 scikit-learn 模型只认 numpy.ndarray 或 pandas.DataFrame,且列顺序、缺失值处理、类别编码都必须对齐,否则报 ValueError: X has 5 features, but RandomForest expected 7 这类错。
最稳妥的方案是定义 Pydantic 模型做校验 + 显式转换:
from pydantic import BaseModel
import numpy as np
<p>class PredictionRequest(BaseModel):
sepal_length: float
sepal_width: float
petal_length: float
petal_width: float</p><p>@app.post("/predict")
def predict(req: PredictionRequest):</p><h1>严格按训练时的列序构造 array</h1><pre class="brush:php;toolbar:false;">X = np.array([[req.sepal_length, req.sepal_width,
req.petal_length, req.petal_width]])
y_pred = model.predict(X)
return {"prediction": int(y_pred[0])}- 不要用
dict(req).values()构造数组——字典无序,字段顺序不保证 - 如果训练时用了
OneHotEncoder或StandardScaler,这些预处理器也得一起加载并串在推理链路里 - 整数特征传了小数(比如
"age": 25.0)一般不影响,但scikit-learn对int64和float64敏感,建议统一转float32避免隐式转换失败
并发请求下 model.predict() 本身是线程安全的,但要注意状态污染
sklearn 大多数 estimator 的 predict 方法是纯函数式、无内部状态的,多线程/异步并发调用没问题。真正要防的是你自己加的逻辑:比如缓存中间结果、修改模型属性、或在预测中调用了带副作用的自定义函数。
- 别在
predict里写model.classes_ = [...]—— 这会污染其他请求看到的模型状态 - 如果用了
CalibratedClassifierCV并启用了n_jobs > 1,注意它底层用joblib.Parallel,可能和 FastAPI 的异步事件循环冲突,建议设n_jobs=1 - 异步包装
model.predict()没意义——它本身不阻塞,强行用run_in_executor反而增加调度开销
模型输出要主动降维,别让 FastAPI 自动 JSON 化 numpy 类型
直接 return {"prob": model.predict_proba(X)[0]} 会触发 FastAPI 尝试序列化 numpy.ndarray,抛出 TypeError: Object of type ndarray is not JSON serializable。
必须显式转成原生 Python 类型:
y_proba = model.predict_proba(X)[0].tolist() # → list[float]
y_pred = int(model.predict(X)[0]) # → int
return {"prediction": y_pred, "confidence": y_proba}
-
.tolist()是必须的,json.dumps(np.array([1,2]))会失败,但json.dumps(np.array([1,2]).tolist())可以 - 如果输出是
np.int64,直接int(x)转,别用int(float(x))多此一举 - 别依赖
orjson或ujson自动支持 numpy——FastAPI 默认不用它们,且行为不一致,显式转换最稳
模型热更新、特征在线验证、A/B 测试路由这些进阶需求,得靠额外服务编排,不是加几行 joblib.load 就能解决的。先确保单模型路径干净、可测、可压测,再谈扩展。
大量免费API接口:立即使用
涵盖生活服务API、金融科技API、企业工商API、等相关的API接口服务。免费API接口可安全、合规地连接上下游,为数据API应用能力赋能!











