我正在使用flask、tensorflow和keras模型构建一个多线程restapi。得到之后
this
错误,我做了一些研究,得出了以下解决方案:
executor = ThreadPoolExecutor(10)
@app.route('/createLearningTask', methods=['POST'])
def createLearningTask():
request_data = request.get_json(force=True)
executor.submit(LearningTask().processData)
return ('', 200)
基本上,我会将每个新的职位申请提交给执行者。在每个请求中,我用给定的参数建立模型,生成模型,预测并存储另一个GET请求的结果。
class LearningTask:
resultDict = {} # access to this map is protected by locks, which I omitted in here
def processData(self, **kwargs):
graph = tf.Graph() # tf = tensorflow
with graph.as_default():
with tf.Session().as_default():
model = Sequential() # keras model
model.add(..)
model.add(..)
model.add(..)
model.compile(..)
model.fit(..)
model.predict(..)
我省略了代码中不相关的部分。它工作得很好,在处理数据并得到结果之后,我把它保存到字典里。我为每个post请求创建新的图表和新的会话。
读后
this
和
this
在讨论中,我提出了我的解决方案。
我的问题是,这个解决方案是安全的和正确的使用张量流?在
documentation
,它说graph不是线程安全的,但是我做了一些负载测试,它可以毫无问题地处理并发请求。