【发布时间】:2018-10-18 10:25:46
【问题描述】:
我有以下问题。我正在使用 Tensorflow Keras 模型来评估连续传感器数据。我的模型输入由 15 个传感器数据帧组成。因为函数 model.predict() 需要将近 1 秒,所以我想异步执行这个函数,这样我就可以在这个时间段内收集下一个数据帧。 为此,我创建了一个带有多处理库和一个用于 model.predict 的函数的池。我的代码如下所示:
def predictData(data):
return model.predict(data)
global model
model = tf.keras.models.load_model("Network.h5")
model._make_predict_function()
p = Pool(processes = 4)
...
res = p.apply_async(predictData, ([[iinput]],))
print(res.get(timeout = 10))
现在我在调用 predictData() 时总是遇到超时错误。似乎 model.predict() 无法正常工作。我做错了什么?
【问题讨论】:
-
tensorflow后端构建的计算图存在于python框架之外。以这种方式使用多处理不会构建图的多个副本。您仍然只有一个模型副本,并试图一次发送 4 个数据流。
-
James 是对的,考虑在前台进程运行预测并在后台运行线程/进程以收集您的下一个数据帧。您可以缓冲多个数据帧,将它们放在网络输入的批处理维度中
-
okey 然后通过线程收集数据似乎是正确的方法。所以一般来说不可能运行例如多个进程中的多个预测?
标签: python tensorflow keras multiprocessing