如何在Flask中跨请求复用Sklearn Pickle模型以优化内存占用
解决方案:Flask中跨请求共享大Sklearn模型的几种方案
嘿,这个问题我太熟了——700MB的Pickle模型每次请求加载确实会把内存搞崩,还严重拖慢响应速度。完全可以在Flask里实现跨请求复用模型,甚至比你想的更简单,下面给你捋几个靠谱的方案:
方案一:全局加载模型(最直接高效的方式)
这是最常用的做法:在Flask应用启动时就把模型加载到内存里,所有请求直接复用这个全局实例。
代码示例
from flask import Flask, request import pickle app = Flask(__name__) # 应用启动时一次性加载模型,全局变量供所有请求使用 with open("large_model.pkl", "rb") as f: model = pickle.load(f) @app.route("/predict", methods=["POST"]) def predict(): data = request.get_json() # 直接用全局的model执行预测,无需重复加载 result = model.predict(data["features"]) return {"prediction": result.tolist()}
注意事项
- 如果用多进程WSGI服务器(比如Gunicorn开多个worker),每个worker进程会独立加载一次模型,内存占用是
worker数量 × 700MB,要根据服务器内存合理调整worker数量。 - 绝对不要在请求处理函数里加载模型,那会导致每次请求都重新读文件、反序列化,完全失去优化意义。
方案二:用Flask-Caching缓存模型(支持动态重载)
如果你的模型需要定期更新,不想每次更新都重启服务,可以用Flask-Caching把模型缓存起来,既实现复用,又能随时刷新。
代码示例
from flask import Flask, request from flask_caching import Cache import pickle app = Flask(__name__) # 配置缓存:用进程内缓存,永不过期(0表示无超时) cache = Cache(app, config={ "CACHE_TYPE": "SimpleCache", "CACHE_DEFAULT_TIMEOUT": 0 }) @cache.memoize() def load_model(): # 这个函数只会被执行一次,结果会被缓存 with open("large_model.pkl", "rb") as f: return pickle.load(f) @app.route("/predict", methods=["POST"]) def predict(): model = load_model() # 直接从缓存取,无需重复加载 data = request.get_json() result = model.predict(data["features"]) return {"prediction": result.tolist()} # 新增接口:手动刷新模型缓存,无需重启服务 @app.route("/refresh-model") def refresh_model(): cache.delete_memoized(load_model) return {"status": "模型已刷新,下次请求会加载新模型"}
注意事项
SimpleCache是进程内缓存,多进程模式下每个worker还是会有自己的缓存副本,和全局加载的内存占用差不多,但多了动态刷新的能力。- 如果要跨进程共享缓存(比如所有worker共用一份模型),可以换成Redis或Memcached作为缓存后端,但700MB的模型存Redis需要确保你的Redis有足够内存,且网络传输不会成为瓶颈。
方案三:共享内存(极端内存紧张场景)
如果服务器内存有限,多进程加载模型会导致OOM,可以用Python的共享内存机制,让所有进程共享同一份模型数据,大幅降低内存占用。
代码思路示例
from flask import Flask, request import pickle import multiprocessing as mp import os import io app = Flask(__name__) shared_model = None def init_shared_memory(): # 第一次启动时,把模型加载到共享内存 with open("large_model.pkl", "rb") as f: model_bytes = f.read() # 创建共享内存区域 shm = mp.SharedMemory(create=True, size=len(model_bytes)) shm.buf[:len(model_bytes)] = model_bytes # 保存共享内存名称,供其他进程连接 with open("shm_name.txt", "w") as f: f.write(shm.name) # 反序列化模型 buf = io.BytesIO(shm.buf[:len(model_bytes)]) return pickle.load(buf) def connect_shared_memory(): # 已有共享内存时,直接连接并加载模型 with open("shm_name.txt", "r") as f: shm_name = f.read() shm = mp.SharedMemory(name=shm_name) buf = io.BytesIO(shm.buf) return pickle.load(buf) @app.before_first_request def setup_model(): global shared_model if os.path.exists("shm_name.txt"): shared_model = connect_shared_memory() else: shared_model = init_shared_memory() @app.route("/predict", methods=["POST"]) def predict(): data = request.get_json() result = shared_model.predict(data["features"]) return {"prediction": result.tolist()}
注意事项
- 这个实现相对复杂,需要处理共享内存的创建、连接和清理,避免内存泄漏。
- 适合内存非常紧张的场景,比如服务器只有2GB内存,多进程加载模型会直接OOM,用共享内存只占700MB左右。
额外优化建议:从根源减少模型内存占用
除了共享模型,还可以优化模型本身,从源头降低内存压力:
- 用
joblib代替pickle保存Sklearn模型:joblib对Sklearn的numpy数组支持更好,生成的文件更小,加载速度也更快。 - 模型剪枝/量化:比如对随机森林、XGBoost等集成模型,减少树的数量或深度;或者用Sklearn的量化工具压缩模型,在精度可接受的前提下缩小体积。
- 拆分大模型:如果模型是多个子模型的组合,可以拆分后分别加载,按需调用。
最后提一句:从IIS迁移到Flask的话,推荐用Gunicorn或uWSGI作为WSGI服务器,前端搭配Nginx反向代理,比直接用IIS部署Flask更高效稳定。
内容的提问来源于stack exchange,提问作者u1234
相关产品推荐
相关产品推荐

