【问题标题】:How to fix keras.backend.eval being really slow如何修复 keras.backend.eval 非常慢
【发布时间】:2019-06-06 17:06:13
【问题描述】:

我想在训练后评估每个样本的模型损失函数。只调用 loss 会为每个批次生成一个值,因此我在 predict()ed 值上手动调用损失函数。

这需要一个张量评估,因为损失返回一个张量。评估这个张量很容易,但调用需要很长时间,尽管是一个简单的操作。

我已经尝试过 keras 会话中的 session.run 以及 keras.backend.eval,两者都有相同的问题。我也尝试过升级 keras,但它已经在 2.2.4 上

import keras
indim = 28
model = Sequential([Dense(8,input_shape=(indim,),activation='tanh'),Dense(4,activation='tanh'),Dense(1,activation='linear')])
model.compile(optimizer='adam',loss='mae')
def foo():
    for i in range(0,500):
        input = np.random.rand(32,28)
        Y     = np.random.rand(32,1)
        Ypred = model.predict(input)
        loss = model.loss_functions[0](Y,Ypred)
        loss = keras.backend.eval(loss)
%prun foo()

我希望上面的例子能在几分之一秒内完成。第一次需要 20 秒,第二次运行需要 40 秒,分析器返回:

      500   27.580    0.055   27.580    0.055 {built-in method _pywrap_tensorflow_internal.ExtendSession}
      500   18.866    0.038   18.866    0.038 {built-in method _pywrap_tensorflow_internal.TF_SessionRun_wrapper}
    16500    0.124    0.000    0.129    0.000 pywrap_tensorflow_internal.py:39(_swig_setattr_nondynamic)

后续调用需要的时间越来越长(20、40、80 秒!)

【问题讨论】:

  • 你只想丢失每批吗? model.evaluate呢?
  • 我想要丢失每个样本。当然,我可以制作 1 大小的批次并对其进行评估,但我认为这会更慢,并且预计 .predict 调用(通常)需要更长的时间

标签: python tensorflow keras


【解决方案1】:

解决方案最终是使用 K.placeholder。否则,全局图会随着模型损失函数的每次调用而增长。

import time

import numpy as np

import keras.backend as K
from keras.layers import Dense
from keras.models import Sequential

indim = 28
model = Sequential([Dense(8, input_shape=(indim,), activation='tanh'), Dense(4, activation='tanh'),
                    Dense(1, activation='linear')])
model.compile(optimizer='adam', loss='mae')
y_pred = K.placeholder([None, 1])
y_true = K.placeholder([None, 1])
loss_fn = model.loss_functions[0](y_true, y_pred)

for i in range(0, 500):
    s = time.time()
    input = np.random.rand(32, 28)
    Y = np.random.rand(32, 1)
    Ypred = model.predict(input)
    _ = K.get_session().run(loss_fn, feed_dict={y_true: Y, y_pred: Ypred})
    print("Took", time.time() - s)

【讨论】:

    猜你喜欢
    • 1970-01-01
    • 1970-01-01
    • 1970-01-01
    • 2015-05-22
    • 1970-01-01
    • 2016-11-18
    • 2017-12-12
    • 2015-06-14
    • 2013-05-14
    相关资源
    最近更新 更多