You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

不使用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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.08.01 02:50:27