Python Flask调用TensorFlow Keras加载模型预测时内存泄漏问题求助
问题:Flask集成TensorFlow模型预测时内存持续增长,疑似内存泄漏
我开发了一个Flask应用,在使用tensorflow.keras.models.load_model进行预测时出现内存泄漏,无法定位原因。该Web API每100ms被调用一次,内存随时间持续增长。我已查阅相关讨论帖,但未找到可行解决方案。
以下是我的Flask应用代码:
#%% import import numpy as np from tensorflow.keras.models import load_model from flask import Flask, request from tensorflow import convert_to_tensor from tensorflow.keras import backend as K import gc APP = Flask(__name__) #%% Function @APP.route('/app') def app(): ## load model model = load_model('model.h5', compile = False) ## input parameter p1 = float(request.args.get('p1')).astype('float32') p2 = float(request.args.get('p2')).astype('float32') ## combine parameter p = np.zeros((1, 2)) p[0,1] = p1 p[0,2] = p2 ## predicted predicted = model(convert_to_tensor(p)) K.clear_session() ## clear gc.collect() model = None del model, p return str(predicted) #%% main if __name__ == "__main__": APP.run(debug = False, host = '0.0.0.0', port = 12345)
环境信息:
- Python 3.10
- Flask 2.3.3
- tensorflow 2.13.0
请问这确实是内存泄漏问题吗?有哪些额外修复措施可采取?
回答
这确实是内存泄漏问题,核心原因分析:
你的代码里每次请求都重新加载一次模型,这是导致内存爆炸的主要元凶。TensorFlow加载模型时会在GPU/CPU内存中分配大量资源(权重张量、计算图等),即使调用K.clear_session()和手动删除变量,也无法完全清理所有残留内存占用,尤其是在100ms一次的高频请求场景下,内存回收速度远赶不上资源分配速度,最终导致内存持续增长。
具体修复措施:
全局加载模型,避免每次请求重复加载
把模型加载代码放在Flask路由函数外面,仅在应用启动时加载一次,所有请求复用同一个模型实例:#%% import import numpy as np import tensorflow as tf from tensorflow.keras.models import load_model from flask import Flask, request from tensorflow import convert_to_tensor import gc APP = Flask(__name__) # 配置TensorFlow按需分配内存(可选,针对GPU场景) gpus = tf.config.list_physical_devices('GPU') if gpus: try: for gpu in gpus: tf.config.experimental.set_memory_growth(gpu, True) except RuntimeError as e: print(e) # 全局加载模型,仅启动时执行一次 model = load_model('model.h5', compile = False) #%% Function @APP.route('/app') def app(): ## input parameter p1 = np.float32(request.args.get('p1')) p2 = np.float32(request.args.get('p2')) ## combine parameter p = np.zeros((1, 2), dtype=np.float32) p[0, 0] = p1 # 修正原代码索引越界问题(数组(1,2)的索引范围是0和1) p[0, 1] = p2 ## predicted predicted = model(convert_to_tensor(p)) # 清理临时变量 del p return str(predicted.numpy()) # 返回numpy数组而非Tensor对象,避免序列化问题修正代码中的索引错误
原代码里p = np.zeros((1, 2))是1行2列的数组,索引范围为[0,0]和[0,1],但你写了p[0,2] = p2,会触发索引越界错误,必须修正。优化内存管理逻辑
- 全局模型复用场景下,不要调用
K.clear_session(),否则会销毁已加载的模型实例,导致后续请求报错。 - 无需频繁手动调用
gc.collect(),Python垃圾回收机制会自动处理临时变量,频繁调用反而增加性能开销。
- 全局模型复用场景下,不要调用
切换到生产级WSGI服务器
Flask自带的APP.run()是开发服务器,不适合高频请求场景。建议使用Gunicorn、uWSGI等生产级服务器,配合合理的多进程/多线程配置,更好地管理内存资源。
内容的提问来源于stack exchange,提问作者黃寶民
相关产品推荐
相关产品推荐

