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

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.05.26 11:10:24