在Flask应用中使用Keras模型时的多线程问题求助
解决Flask单线程高并发瓶颈+模型重复加载问题
你现在面临两个核心痛点:一是Flask默认单线程扛不住每秒100+的请求,还没法用threaded=True;二是每次请求都重新加载模型,直接把响应速度拖垮了。下面给你一套落地的解决方案:
第一步:提前加载所有模型,彻底避免请求时重复加载
每次请求都读取structure.json和加载weight.h5是极大的资源浪费——模型加载是IO密集+计算密集的操作,绝对不能放在请求处理逻辑里。你应该在Flask应用启动时就把三个模型全部加载到内存里,后续所有请求直接复用:
from flask import Flask, request, jsonify from keras.models import model_from_json api = Flask(__name__) # 封装模型加载的通用函数 def load_keras_model(structure_path, weights_path): # 读取模型结构 with open(structure_path, "r") as f: model_structure = f.read() model = model_from_json(model_structure) # 加载预训练权重 model.load_weights(weights_path) return model # 启动应用时一次性加载三个模型,存在全局变量中 # 注意:请根据你的实际文件路径调整 model_a = load_keras_model("./models/model_a/structure.json", "./models/model_a/weight.h5") model_b = load_keras_model("./models/model_b/structure.json", "./models/model_b/weight.h5") model_c = load_keras_model("./models/model_c/structure.json", "./models/model_c/weight.h5") # 你的GET接口示例 @api.route("/health", methods=["GET"]) def health_check(): return jsonify({"status": "ok"}), 200 # 你的POST接口示例(直接复用已加载的模型) @api.route("/predict", methods=["POST"]) def predict(): data = request.get_json() # 根据业务逻辑选择对应模型推理 if data.get("model_type") == "a": result = model_a.predict(data["input_data"]) elif data.get("model_type") == "b": result = model_b.predict(data["input_data"]) else: result = model_c.predict(data["input_data"]) return jsonify({"prediction": result.tolist()}), 200 # 保留基础的run()即可,无需添加任何参数 if __name__ == "__main__": api.run()
这样修改后,每个请求进来直接用已经加载好的模型,响应速度会提升一个数量级。
第二步:替换Flask自带服务器,用生产级WSGI服务器扛高并发
Flask自带的api.run()只是给开发调试用的“玩具服务器”,哪怕开了threaded=True也扛不住生产环境的高并发。既然你没法用threaded=True,直接换用Gunicorn(最易用的生产级WSGI服务器),它可以通过命令行参数轻松配置多进程+多线程,完全绕开Flask的run参数限制。
安装Gunicorn
pip install gunicorn
启动应用(配置多进程多线程)
假设你的Flask代码文件叫app.py,应用实例是api,执行这条命令:
gunicorn --workers=4 --threads=2 --bind=0.0.0.0:5000 app:api
参数说明:
--workers:进程数,建议设为你CPU核心数的1-2倍(比如4核CPU设4或8)--threads:每个进程的线程数,建议设2-4--bind:绑定的IP和端口,确保外部能正常访问app:api:指定你的Flask应用所在文件(app.py)和应用实例名(api)
如果你的模型推理是CPU密集型,可以适当减少线程数、增加进程数;如果是GPU推理,注意不要开太多进程导致GPU内存不足,根据你的GPU显存调整workers数量。
额外优化建议
- 模型轻量化:如果模型体积过大,可以考虑用TensorFlow Lite或者ONNX Runtime进行模型优化,减少内存占用和推理时间
- 请求排队缓冲:如果请求量持续超过服务器处理能力,可以在前面加一层Nginx做反向代理和请求排队,避免直接压垮应用
- 高频请求缓存:如果存在重复的请求输入,可以加个Redis缓存,直接返回之前的推理结果,减少模型计算量
这样一套操作下来,处理每秒100+的请求完全没问题。
内容的提问来源于stack exchange,提问作者Le Duong Tuan Anh
相关产品推荐
相关产品推荐

