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

如何在tf.keras中获取每个train_step的各阶段详细运行时长?

在tf.keras中获取train_step各子步骤的运行时长

要在model.fit()中精准追踪每个训练步骤里的细分操作耗时(预测、损失计算、梯度计算、梯度应用、指标更新),最直接的方式是自定义Model类并重写train_step方法,把计时逻辑嵌入到每个子步骤中;另外也可以用TensorFlow Profiler做更深度的性能分析。下面详细说明两种方案:

方案一:自定义Model重写train_step(推荐,可实时获取数值)

这种方式完全兼容model.fit()的使用习惯,同时能精准统计每个子步骤的耗时,还能通过回调实时打印结果。

1. 定义带计时功能的Model类

我们会在train_step中嵌入tf.timestamp()(图模式和eager模式都兼容)来统计每个阶段的时间,并用tf.keras.metrics.Mean来记录每个step的平均耗时:

import tensorflow as tf
from tensorflow.keras import Model, layers, optimizers, metrics

class TimedTrainingModel(Model):
    def __init__(self, *args, **kwargs):
        super().__init__(*args, **kwargs)
        # 初始化训练精度指标
        self.train_accuracy = metrics.SparseCategoricalAccuracy(name="train_acc")
        # 初始化各个时间统计指标
        self.prediction_time = metrics.Mean(name="prediction_time")
        self.loss_calculate_time = metrics.Mean(name="loss_calculate_time")
        self.grad_compute_time = metrics.Mean(name="grad_compute_time")
        self.grad_apply_time = metrics.Mean(name="grad_apply_time")
        self.metrics_update_time = metrics.Mean(name="metrics_update_time")

    def train_step(self, inputs):
        images, labels = inputs

        # 1. 统计预测阶段耗时
        start_ts = tf.timestamp()
        with tf.GradientTape() as tape:
            predictions = self(images, training=True)
        pred_end_ts = tf.timestamp()
        self.prediction_time.update_state(pred_end_ts - start_ts)

        # 2. 统计损失计算耗时(兼容compile时配置的损失和正则化)
        start_ts = tf.timestamp()
        loss = self.compiled_loss(
            labels, predictions, regularization_losses=self.losses
        )
        loss_end_ts = tf.timestamp()
        self.loss_calculate_time.update_state(loss_end_ts - start_ts)

        # 3. 统计梯度计算耗时
        start_ts = tf.timestamp()
        gradients = tape.gradient(loss, self.trainable_variables)
        grad_end_ts = tf.timestamp()
        self.grad_compute_time.update_state(grad_end_ts - start_ts)

        # 4. 统计梯度应用耗时(兼容compile时配置的优化器)
        start_ts = tf.timestamp()
        self.compiled_optimizer.apply_gradients(
            zip(gradients, self.trainable_variables)
        )
        apply_end_ts = tf.timestamp()
        self.grad_apply_time.update_state(apply_end_ts - start_ts)

        # 5. 统计指标更新耗时
        start_ts = tf.timestamp()
        self.train_accuracy.update_state(labels, predictions)
        # 同步更新compile时配置的其他指标
        self.compiled_metrics.update_state(labels, predictions)
        update_end_ts = tf.timestamp()
        self.metrics_update_time.update_state(update_end_ts - start_ts)

        # 返回所有指标结果,会在fit的日志中显示
        return {m.name: m.result() for m in self.metrics}

2. 使用自定义Model并添加回调实时打印

接下来就可以像普通tf.keras模型一样编译、训练,还能写个回调函数来定期打印每个step的细分耗时:

# 构建示例模型(这里用MNIST分类任务为例)
model = TimedTrainingModel(
    layers.Conv2D(32, (3,3), activation="relu", input_shape=(28,28,1)),
    layers.MaxPooling2D(),
    layers.Flatten(),
    layers.Dense(10, activation="softmax")
)

# 编译模型(和普通模型完全一致)
model.compile(
    optimizer=optimizers.Adam(learning_rate=1e-3),
    loss=tf.keras.losses.SparseCategoricalCrossentropy()
)

# 加载并预处理数据
(x_train, y_train), _ = tf.keras.datasets.mnist.load_data()
x_train = x_train[..., tf.newaxis].astype("float32") / 255.0

# 自定义回调:每100个step打印一次耗时统计
class StepTimeLogger(tf.keras.callbacks.Callback):
    def on_train_batch_end(self, batch, logs=None):
        if batch % 100 == 0:
            tf.print(f"\n=== Batch {batch} Time Stats ===")
            tf.print(f"Prediction: {self.model.prediction_time.result():.6f}s")
            tf.print(f"Loss Calculation: {self.model.loss_calculate_time.result():.6f}s")
            tf.print(f"Gradient Compute: {self.model.grad_compute_time.result():.6f}s")
            tf.print(f"Gradient Apply: {self.model.grad_apply_time.result():.6f}s")
            tf.print(f"Metrics Update: {self.model.metrics_update_time.result():.6f}s")
            # 重置指标,让下一个区间的统计是新的均值
            self.model.prediction_time.reset_states()
            self.model.loss_calculate_time.reset_states()
            self.model.grad_compute_time.reset_states()
            self.model.grad_apply_time.reset_states()
            self.model.metrics_update_time.reset_states()

# 启动训练
model.fit(
    x_train, y_train,
    epochs=3,
    batch_size=32,
    callbacks=[StepTimeLogger()],
    verbose=1
)

方案二:用TensorFlow Profiler做深度性能分析

如果你需要更底层的操作耗时(比如GPU内核执行时间、内存占用等),可以用TensorFlow Profiler:

# 启动Profiler服务(端口可自定义)
tf.profiler.experimental.server.start(6009)

# 运行一轮训练
model.fit(x_train, y_train, epochs=1, batch_size=32)

# 停止Profiler服务
tf.profiler.experimental.server.stop()

然后打开TensorBoard,切换到Profile标签页,就能看到所有操作的耗时分布,包括你关心的预测、梯度计算等阶段的底层操作耗时。不过这种方式更适合性能调优,而非实时获取每个step的数值统计。

注意事项

  • 必须用tf.timestamp()而非Python的time.time(),前者在TensorFlow图模式下能正常工作,后者会被固化到图中导致时间不更新。
  • 重写train_step时要调用self.compiled_loss和self.compiled_optimizer,这样能保持和model.compile()配置的一致性(比如正则化损失、优化器参数)。

内容的提问来源于stack exchange,提问作者Jiayuan Li

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.05.06 17:22:48