部署TensorFlow Serving模型时遇“Attempted to use a Closed Session”错误求助
解决Flask+TensorFlow Serving调用中"Attempted to use a Closed Session"错误
嗨,我来帮你搞定这个烦人的错误!这个问题本质上是TensorFlow Session的生命周期和Flask的请求处理流程不匹配导致的,结合你提到的要调用TensorFlow Serving模型的需求,我给你两个针对性的解决方案:
方案1:如果是本地加载模型(非远程Serving),正确管理Session生命周期
如果你的代码是本地训练后直接加载模型,没有用TensorFlow Serving的话,问题大概率出在Session被提前关闭或者跨请求复用了已关闭的Session。Flask是多线程框架,全局Session很容易出问题,你可以这样调整:
- 在Flask应用启动时初始化Session并存在app配置里,确保它在整个应用生命周期内保持活跃:
from flask import Flask import tensorflow as tf import numpy as np import os app = Flask(__name__) DATA_SIZE = 100 SAVE_PATH = './save' MODEL_NAME = 'test' # 应用启动时加载模型和Session with tf.Session() as sess: saver = tf.train.import_meta_graph(os.path.join(SAVE_PATH, f"{MODEL_NAME}.meta")) saver.restore(sess, tf.train.latest_checkpoint(SAVE_PATH)) # 把Session和Graph存到app配置,供请求使用 app.config['TF_SESSION'] = sess app.config['TF_GRAPH'] = tf.get_default_graph() @app.route('/predict', methods=['POST']) def predict(): # 从app配置取出活跃的Session和Graph sess = app.config['TF_SESSION'] graph = app.config['TF_GRAPH'] # 必须在当前Graph的上下文中执行操作 with graph.as_default(): # 这里替换成你的实际预测逻辑,比如获取输入数据 input_data = np.random.rand(1, DATA_SIZE) # 注意替换成你模型中实际的张量名称 prediction = sess.run('prediction_op:0', feed_dict={'input_tensor:0': input_data}) return str(prediction) if __name__ == '__main__': app.run()
如果你用的是TensorFlow 2.x,建议直接用Keras的
load_model加载模型,TF2.x默认是即时执行模式,不需要手动管理Session,会省心很多。
方案2:通过HTTP/gRPC调用TensorFlow Serving(更符合你的需求)
既然你明确要调用网络可访问的TensorFlow Serving模型,那根本不需要在Flask里维护TensorFlow Session!直接通过HTTP请求转发到Serving服务就好,这才是正确的架构方式:
HTTP调用示例代码
from flask import Flask, request, jsonify import requests import numpy as np app = Flask(__name__) # 替换成你的TensorFlow Serving地址,默认端口8501是HTTP端口,8500是gRPC端口 TF_SERVING_ENDPOINT = 'http://localhost:8501/v1/models/test:predict' @app.route('/predict', methods=['POST']) def predict(): # 获取前端传来的输入数据 input_data = request.json.get('data') if not input_data: return jsonify({'error': '请提供输入数据'}), 400 # 构造TensorFlow Serving要求的JSON格式 serving_request = { 'instances': [input_data] } # 发送请求到TensorFlow Serving resp = requests.post(TF_SERVING_ENDPOINT, json=serving_request) if resp.status_code != 200: return jsonify({'error': '调用TensorFlow Serving失败'}), resp.status_code # 解析返回的预测结果 predictions = resp.json().get('predictions') return jsonify({'predictions': predictions}) if __name__ == '__main__': app.run(host='0.0.0.0', port=5000)
为什么这个方案更好?
- TensorFlow Serving本身就是专门用来管理模型生命周期、处理推理请求的服务,它会自动处理Session、Graph的创建和维护,你完全不用操心。
- Flask只需要作为API网关,负责接收用户请求、转发到Serving、返回结果,逻辑更清晰,也彻底避免了Session相关的错误。
最后再给你几个排查小技巧
- 检查代码里有没有
session.close()的调用,如果有,确保它不会在请求处理完成前执行。 - 如果用多进程/多线程的Flask服务器(比如gunicorn),别在全局创建Session,每个进程/线程的上下文是独立的,要么每个请求重新创建,要么用线程本地存储。
- 对于TF1.x,一定要确保执行操作时,当前的Graph是默认Graph,并且Session处于活跃状态。
内容的提问来源于stack exchange,提问作者Shravan Kumar
相关产品推荐
相关产品推荐

