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

如何在Keras及TensorFlow2中获取模型平均训练速度

TensorFlow 2 对应记录global_step/sec指标的实现方式

在TensorFlow 2中,根据你使用的训练流程不同,对应实现方式分为两类:

1. 使用Keras model.fit() 高阶训练API

这种场景直接使用内置的 tf.keras.callbacks.TensorBoard 回调即可实现对应功能:

  • 回调默认会自动生成global_step/sec指标写入tfevent文件,不需要额外配置
  • 参数update_freq对应原TF1 Estimator的log_step_count_steps,传入整数N就代表每N个训练步记录一次指标
  • 示例代码:
import tensorflow as tf

# 定义TensorBoard回调
tensorboard_cb = tf.keras.callbacks.TensorBoard(
    log_dir="./train_logs",
    update_freq=10, # 每10步记录一次步速指标
    profile_batch=0 # 无需性能分析时可关闭,降低额外开销
)

# 训练时传入回调即可
model.fit(train_dataset, epochs=20, callbacks=[tensorboard_cb])

2. 使用自定义训练循环

如果是自己写的训练循环,需要手动统计步速并写入tfevent文件,示例逻辑如下:

import time
import tensorflow as tf

# 初始化日志写入器
log_writer = tf.summary.create_file_writer("./train_logs")
# 配置记录步长,对应原log_step_count_steps
LOG_STEP_FREQ = 10

global_step = 0
last_record_time = time.time()
last_record_step = 0

# 训练循环
for epoch in range(20):
    for batch in train_dataset:
        global_step += 1
        # 你的训练步逻辑,例如:
        with tf.GradientTape() as tape:
            pred = model(batch[0], training=True)
            loss = loss_fn(batch[1], pred)
        grads = tape.gradient(loss, model.trainable_variables)
        optimizer.apply_gradients(zip(grads, model.trainable_variables))

        # 到指定步长时统计步速并写入日志
        if global_step % LOG_STEP_FREQ == 0:
            current_time = time.time()
            step_delta = global_step - last_record_step
            time_delta = current_time - last_record_time
            steps_per_sec = step_delta / time_delta
            with log_writer.as_default():
                tf.summary.scalar("global_step/sec", steps_per_sec, step=global_step)
            # 更新记录锚点
            last_record_time = current_time
            last_record_step = global_step

两种方式生成的tfevent文件格式和TensorFlow 1完全兼容,你可以沿用之前提取global_step/sec指标计算平均训练速度的逻辑。

内容的提问来源于stack exchange,提问作者Muyang Yu

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.10.03 12:06:00