不使用model.fit()时如何关闭Keras训练轮次日志
关闭GradientTape训练DQN时的进度日志输出
问题描述
我目前正在用GradientTape训练深度Q网络(DQN),代码如下:
with tf.GradientTape() as tape: q_values_current_state_dqn = self.dqn_architecture(states) one_hot_actions = tf.keras.utils.to_categorical(actions, self.num_legal_actions, dtype=np.float32) # e.g. [[0,0,1,0],[1,0,0,0],...] q_values_current_state_dqn = tf.reduce_sum(tf.multiply(q_values_current_state_dqn, one_hot_actions), axis=1) error = q_values_current_state_dqn - target_q_values loss = tf.keras.losses.Huber()(target_q_values, q_values_current_state_dqn) dqn_architecture_gradients = tape.gradient(loss, self.dqn_architecture.trainable_variables) # Computes the gradient using operations recorded in context of this tape. self.dqn_architecture.optimizer.apply_gradients(zip(dqn_architecture_gradients, self.dqn_architecture.trainable_variables))
但训练时会不断输出如下进度日志,希望关闭:
1/1 [==============================] - 0s 34ms/step 1/1 [==============================] - 0s 10ms/step 1/1 [==============================] - 0s 11ms/step 1/1 [==============================] - 0s 10ms/step 1/1 [==============================] - 0s 11ms/step 1/1 [==============================] - 0s 10ms/step 1/1 [==============================] - 0s 13ms/step 1/1 [==============================] - 0s 10ms/step 1/1 [==============================] - 0s 10ms/step 1/1 [==============================] - 0s 10ms/step 1/1 [==============================] - 0s 11ms/step 1/1 [==============================] - 0s 10ms/step 1/1 [==============================] - 0s 11ms/step 1/1 [==============================] - 0s 10ms/step 1/1 [==============================] - 0s 10ms/step 1/1 [==============================] - 0s 10ms/step 1/1 [==============================] - 0s 10ms/step 1/1 [==============================] - 0s 12ms/step
我知道用model.fit()时可以设置verbose=0关闭日志,但不清楚用GradientTape时该怎么操作。
解决方案
这些日志是Keras默认的交互式训练日志输出,你可以通过以下两种方式关闭:
1. 全局关闭交互式日志
在训练代码的开头添加一行代码,一次性关闭所有Keras的交互式日志输出:
tf.keras.utils.disable_interactive_logging()
这会关闭包括进度条、每步训练日志在内的所有交互式输出,后续所有模型的前向/反向传播都不会再打印这类日志。
2. 临时关闭日志(上下文管理器)
如果你只想在GradientTape训练的这段代码中关闭日志,其他地方保留输出,可以使用上下文管理器临时禁用:
import tensorflow as tf with tf.keras.utils.disable_interactive_logging(): with tf.GradientTape() as tape: q_values_current_state_dqn = self.dqn_architecture(states) one_hot_actions = tf.keras.utils.to_categorical(actions, self.num_legal_actions, dtype=np.float32) q_values_current_state_dqn = tf.reduce_sum(tf.multiply(q_values_current_state_dqn, one_hot_actions), axis=1) error = q_values_current_state_dqn - target_q_values loss = tf.keras.losses.Huber()(target_q_values, q_values_current_state_dqn) dqn_architecture_gradients = tape.gradient(loss, self.dqn_architecture.trainable_variables) self.dqn_architecture.optimizer.apply_gradients(zip(dqn_architecture_gradients, self.dqn_architecture.trainable_variables))
注意:把optimizer.apply_gradients放在外层不影响,因为日志主要来自模型的前向传播调用。
内容的提问来源于stack exchange,提问作者Bryan
相关产品推荐
相关产品推荐

