能用,但必须提前加载模型、禁用训练模式、注意线程安全;flask启动前完成加载并调用model.eval()或model.trainable=false,输入需转float32并严格匹配形状,并发时须加锁。

直接上结论:能用,但必须提前加载模型、禁用训练模式、注意线程安全——否则首次请求会卡死,多并发时预测结果错乱。
Flask启动前必须完成模型加载和model.eval()
TensorFlow/Keras模型不能在每次请求里重复加载,否则每个request都会触发一次磁盘读取+图构建+权重加载,耗时从毫秒级变成秒级。更严重的是,Keras默认启用trainable=True,若没显式调用model.trainable = False或model.compile后不调用fit,某些层(如BatchNormalization、Dropout)在推理时行为不稳定。
-
load_model('model.h5')必须放在if __name__ == '__main__':之前,或封装进全局变量初始化函数 - PyTorch用户必须加
model.eval();Keras用户建议加model.trainable = False并确认tf.keras.backend.set_learning_phase(0) - GPU用户需确保
torch.cuda.is_available()为True后再调用model.cuda(),否则报RuntimeError: Input type (torch.FloatTensor) and weight type (torch.cuda.FloatTensor) should be the same
predict函数里别用model.predict()直接传原始列表
常见错误是把客户端发来的[[0.1, 0.2, ..., 0.9]]直接喂给model.predict(),结果报ValueError: Input 0 is incompatible with layer...——因为Keras期望float32且维度匹配,而JSON解析后是Python list,numpy自动转成float64,且reshape容易漏掉batch维度。
- 务必用
np.array(data, dtype=np.float32)显式指定类型 - 输入形状必须严格匹配模型定义,比如模型输入是
(None, 224, 224, 3),就要reshape(-1, 224, 224, 3),不能只写reshape(1, -1) - 图像类任务要加预处理:
tf.keras.applications.mobilenet_v2.preprocess_input(img_array),否则像素值范围错(0–255 vs 0–1 vs -1–1)会导致预测全错
并发请求下model.predict()不是线程安全的
TensorFlow 2.x默认使用 eager execution,但底层仍共享计算图资源。实测在gunicorn多worker部署时,若没加锁或隔离,两个请求同时调用model.predict()可能返回对方的预测结果,或抛出InvalidArgumentError: You must feed a value for placeholder。
- 最简单方案:用
threading.Lock()包裹预测逻辑,适合QPS - 生产环境推荐改用
flask+tensorflow-serving或FastAPI+uvicorn,后者天然支持异步+多进程隔离 - 绝对不要在
@app.route里重新load_model——模型加载本身不是原子操作,多线程同时执行会竞争文件句柄,导致OSError: Unable to open file
真正麻烦的不是怎么写通,而是模型加载时机、数据类型对齐、并发控制这三点——它们不出现在任何“Hello World”教程里,但上线后第一个报警一定来自这里。
Python免费学习笔记(深入):立即使用
在学习笔记中,你将探索 Python 的核心概念和高级技巧!











