Keras情感分析模型Flask API二次预测张量报错求助
解决Flask部署Keras模型重复预测报错的问题
问题原因分析
你碰到的「无法将feed_dict键解释为张量:张量不属于此图」错误,核心问题出在这两点:
- 每次请求都通过
pickle.load重新加载模型,还在加载后调用tf.reset_default_graph(),这会直接清空当前TensorFlow的默认图,导致后续请求的输入张量和模型绑定的图完全不匹配。 - Flask是多线程运行环境,Keras模型和TensorFlow的计算图是强绑定的,重复加载+重置图的操作会彻底破坏这种绑定关系,只有第一次请求能在正确的图上下文里正常运行。
修复方案
1. 全局加载资源,避免重复初始化
把模型、tokenizer、max_length这些需要复用的资源,在Flask应用启动时就一次性加载完成,不要每次处理请求都重新加载一遍。
2. 维护TensorFlow图上下文
Keras依赖TensorFlow的计算图运行,需要把模型预测的逻辑包裹在对应的图上下文里,避免多线程环境下的图混乱问题。
修改后的完整代码如下:
import flask from flask import Flask, render_template, request import pickle import tensorflow as tf from keras.preprocessing.text import Tokenizer from keras.preprocessing.sequence import pad_sequences # 定义全局变量,用于存储复用资源 app = Flask(__name__) model = None tokenizer_obj = None max_length = None graph = None def load_resources(): """在应用启动时加载所有需要复用的资源""" global model, tokenizer_obj, max_length, graph # 加载模型(如果是用Keras原生save方法保存的,建议用load_model更可靠) model = pickle.load(open("model.pkl","rb")) print("Loaded model from disk") # 加载tokenizer(需要你提前用pickle保存训练好的tokenizer) # tokenizer_obj = pickle.load(open("tokenizer.pkl","rb")) # 设置训练时确定的max_length值 # max_length = 6 # 获取当前模型绑定的TensorFlow图 graph = tf.get_default_graph() def ValuePredictor(to_predict): """预测函数,在指定图上下文里运行""" global model, graph with graph.as_default(): result = model.predict(to_predict) return result @app.route('/') def welcome(): return flask.render_template('welcome.html') @app.route('/result',methods = ['POST']) def result(): if request.method == 'POST': to_predict_list = request.form.to_dict() to_predict_list = list(to_predict_list.values()) test_tokens = tokenizer_obj.texts_to_sequences(to_predict_list) test_pad = pad_sequences(test_tokens, maxlen = max_length, padding= 'post') print(test_pad) result = ValuePredictor(test_pad) return render_template("result.html",prediction=result) if __name__ == '__main__': # 启动应用前先加载所有资源 load_resources() app.run(debug= True, port = 5000)
额外注意事项
- 一定要确保
tokenizer_obj和max_length也全局加载,不要每次请求重新初始化,否则会导致输入序列的处理规则和训练时不一致。 - 彻底删掉
tf.reset_default_graph()这个操作,它会破坏已有的图绑定关系,完全没必要在预测流程里调用。 - 如果你的模型是用Keras的
model.save()方法保存的,建议改用keras.models.load_model()加载,这个方法会自动处理图的绑定问题,比pickle更稳定可靠:from keras.models import load_model model = load_model("model.h5")
这样修改后,模型只在启动时加载一次,所有请求都在同一个图上下文里运行,就能解决重复预测时的报错问题了。
内容的提问来源于stack exchange,提问作者amro_ghoneim
相关产品推荐
相关产品推荐

