【问题标题】:How to disable the warning "tensorflow:Method (on_train_batch_end) is slow compared to the batch update (). Check your callbacks"如何禁用警告“tensorflow:方法(on_train_batch_end)与批量更新()相比很慢。检查你的回调”
【发布时间】:2020-01-15 10:29:27
【问题描述】:

我正在尝试使用 tensorflow 2.0 以 keras 样式实现 Stochastic Weight Averaging (SWA),因此我需要在每一步都更新 SWA 模型权重。我已经编写了一个自定义回调来执行此操作,但我每一步都收到警告。以下是一些细节:

我的自定义回调:


class CustomCallback(tf.keras.callbacks.Callback):
    def __init__(self, valid_data, output_path, swa_alpha=0.99, eval_every=500, eval_batch=16, fold=None):
        self.valid_inputs = valid_data[0]
        self.valid_outputs = valid_data[1]
        self.eval_batch = eval_batch
        self.swa_alpha = swa_alpha
        self.fold = fold
        self.output_path = output_path
        self.rho_value = -1  # record the best rho for report
        self.eval_every = eval_every

    def on_train_begin(self, logs={}):
        self.swa_weights = self.model.get_weights()

    def on_batch_end(self, batch, logs={}):

        # update swa parameters
        alpha = min(1 - 1 / (batch + 1), self.swa_alpha)
        current_weights = self.model.get_weights()
        for i, layer in enumerate(self.model.layers):
            self.swa_weights[i] = alpha * self.swa_weights[i] + (1 - alpha) * current_weights[i]

        # validation
        if batch > 0 and batch % self.eval_every == 0:
            # do validation
            val_pred = self.model.predict(self.valid_inputs, batch_size=self.eval_batch)
            rho_val = compute_spearmanr(self.valid_outputs, val_pred)  # the metric

            # set the swa parameters and do validation
            self.model.set_weights(self.swa_weights)
            swa_val_pred = self.model.predict(self.valid_inputs, batch_size=self.eval_batch)
            swa_rho_val = compute_spearmanr(self.valid_outputs, swa_val_pred)

            # reset the original parameters
            self.model.set_weights(current_weights)

            # check whether to save model and update best rho value
            if rho_val > self.rho_value:
                self.rho_value = rho_val
                self.model.save_weights(f'{self.output_path}/fold-{fold}-best.h5')

        del current_weights
        gc.collect()

输出是这样的:

WARNING:tensorflow:Method (on_train_batch_end) is slow compared to the batch update (11.428264). Check your callbacks.
WARNING:tensorflow:Method (on_train_batch_end) is slow compared to the batch update (11.464315). Check your callbacks.
WARNING:tensorflow:Method (on_train_batch_end) is slow compared to the batch update (11.502968). Check your callbacks.
WARNING:tensorflow:Method (on_train_batch_end) is slow compared to the batch update (11.518413). Check your callbacks.

我每一步都收到警告,这意味着在不运行验证代码的情况下,用于更新 SWA 参数的代码(self.model.get_weights() 和以下 for 循环)已经足够慢了。

我知道更新参数非常慢,因为model.get_weights()model.set_weights() 都会对参数进行深拷贝(根据我的实验,新的 numpy ndarray 的新列表)。

我认为我的 SWA 实现没有任何问题(如果有任何错误,请告诉我),所以我只想禁用警告。

我尝试过的:

  1. 添加代码os.environ["TF_CPP_MIN_LOG_LEVEL"] = "2"os.environ["TF_CPP_MIN_LOG_LEVEL"] = "3" 以禁用WARNING。
  2. model.fit() 中将verbose 设置为20,即model.fit(..., verbose=2, ...)model.fit(..., verbose=0, ...)

两者都不起作用。

有什么想法吗?提前感谢您的帮助!

【问题讨论】:

    标签: python tensorflow keras warnings tensorflow2.0


    【解决方案1】:

    这不是一个非常令人满意的答案,但 TF_CPP_MIN_LOG_LEVEL 不起作用是一个已知问题:TF_CPP_MIN_LOG_LEVEL does not work with TF2.0 dev20190820

    我能够在tensorflow==2.1.0-rc1 上通过一个玩具示例重现您的问题:

    import os
    import time
    os.environ['TF_CPP_MIN_LOG_LEVEL'] = "2"
    import tensorflow as tf
    tf.get_logger().setLevel("WARNING")
    tf.autograph.set_verbosity(2)
    
    print(tf.__version__)
    
    mnist = tf.keras.datasets.mnist
    
    (x_train, y_train), (x_test, y_test) = mnist.load_data()
    x_train, x_test = x_train / 255.0, x_test / 255.0
    
    model = tf.keras.models.Sequential([
      tf.keras.layers.Flatten(input_shape=(28, 28)),
      tf.keras.layers.Dense(128, activation='relu'),
      tf.keras.layers.Dropout(0.2),
      tf.keras.layers.Dense(10, activation='softmax')
    ])
    
    class CustomCallback(tf.keras.callbacks.Callback):
    
      def on_train_batch_end(self, batch, logs=None):
        time.sleep(3)
    
    model.compile(optimizer='adam',
                  loss='sparse_categorical_crossentropy',
                  metrics=['accuracy'])
    
    model.fit(x_train, y_train, epochs=1, callbacks=[CustomCallback()])
    
    2.1.0
    Downloading data from https://storage.googleapis.com/tensorflow/tf-keras-datasets/mnist.npz
    11493376/11490434 [==============================] - 31s 3us/step
    Train on 60000 samples
    WARNING:tensorflow:Method (on_train_batch_end) is slow compared to the batch update (3.002797). Check your callbacks.
       32/60000 [..............................] - ETA: 1:57:38 - loss: 2.4674 - accuracy: 0.0938WARNING:tensorflow:Method (on_train_batch_end) is slow compared to the batch update (3.002938). Check your callbacks.
    ...
    

    标准建议(os.environ['TF_CPP_MIN_LOG_LEVEL']tf.get_logger().setLevel("WARNING")tf.autograph.set_verbosity(2))都不起作用,我怀疑您必须等到上述问题得到解决。

    【讨论】:

    • 您是否尝试过 tf.get_logger().setLevel("ERROR") 而不是警告(实际上您是否想要隐藏警告)?每晚使用 tf 2.20 它可以正常工作。
    猜你喜欢
    • 1970-01-01
    • 1970-01-01
    • 1970-01-01
    • 1970-01-01
    • 2013-08-22
    • 1970-01-01
    • 1970-01-01
    • 2018-07-30
    • 1970-01-01
    相关资源
    最近更新 更多