如何在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
相关产品推荐
相关产品推荐

