TensorFlow模型集成Flask时抛出“非此图元素”错误求助
解决Flask集成TensorFlow句子嵌入时的
TypeError: Cannot interpret feed_dict key as Tensor错误 这个问题我之前在做TensorFlow和Flask集成时也踩过坑,本质是TensorFlow的图(Graph)和会话(Session)在Flask多线程环境下的上下文不匹配导致的。
错误原因分析
你在Flask启动前全局初始化了Embeddings实例,此时它会在TensorFlow的默认图中创建Placeholder和模型节点。但Flask默认是多线程运行的,每个请求线程会使用自己的默认图上下文;同时你在get_sentence_embeddings里每次都新建tf.Session(),这个会话默认绑定的是当前线程的默认图,和你初始化Embeddings时的图不是同一个,所以就会出现"找不到Placeholder"的错误。
解决方案:绑定专属的图和会话
核心思路是让Embeddings实例拥有自己独立的图和会话,并且在请求处理时强制使用这个图的上下文,避免和线程默认图冲突。
修改后的句子嵌入类代码
import tensorflow as tf import tensorflow_hub as hub class Embeddings: def __init__(self): self.embedding_model_url = config_obj.tf_model_url # 1. 创建实例专属的图,避免和全局默认图冲突 self.graph = tf.Graph() with self.graph.as_default(): # 2. 在专属图的上下文中初始化模型、占位符和输出节点 self.embedding_model = hub.Module(self.embedding_model_url) self.messages = tf.placeholder(dtype=tf.string, shape=[None]) self.output = self.embedding_model(self.messages) # 3. 创建并保存会话,一次性完成变量和表的初始化 self.sess = tf.Session(graph=self.graph) self.sess.run([tf.global_variables_initializer(), tf.tables_initializer()]) def get_sentence_embeddings(self, sentence): # 使用实例自己的图和会话执行计算 with self.graph.as_default(): result = self.sess.run(self.output, feed_dict={self.messages: [sentence]}) return result if __name__ == '__main__': sentence = "GoAir is waiving cancellation and change fees for Bhubaneswar, Kolkata and Ranchi flights for travel between May 2 and May 5, the airline said in a statement" tf_object = Embeddings() embeddings = tf_object.get_sentence_embeddings(sentence) print(embeddings)
修改后的Flask集成代码
from sentence_embeddings import Embeddings from flask import Flask, request, Response import json import time import logging app = Flask(__name__) logger = logging.getLogger(__name__) embedding_obj = Embeddings() @app.route('/get-similar-claims', methods=['POST']) def get_similar_claims(): params = request.get_json() claim = params.get("claim", "") num_results = params.get("num_results", 10) t0 = time.time() # 强制使用Embeddings实例的图上下文 with embedding_obj.graph.as_default(): claim_embeddings = embedding_obj.get_sentence_embeddings(claim) logger.info("Time taken to calculate sentence embeddings - {}".format(round(time.time() - t0, 4))) # 注意:numpy数组不能直接序列化,需要转成列表 return Response(json.dumps(claim_embeddings.tolist()), mimetype='application/json') if __name__ == '__main__': app.run('0.0.0.0', 5001)
关键注意点
- 专属图与会话:让每个
Embeddings实例拥有自己的图和会话,避免多线程环境下的图上下文冲突 - 一次性初始化:只在实例创建时初始化一次变量和表,避免每次请求重复初始化浪费资源
- 序列化处理:TensorFlow返回的是numpy数组,需要用
tolist()转成Python列表才能被json.dumps序列化
内容的提问来源于stack exchange,提问作者Anurag Sharma
相关产品推荐
相关产品推荐

