如何在TensorFlow 2.8.0中统计训练与预测阶段的FLOPS?无需外部工具
无需外部工具的TensorFlow 2.8.0 CPU环境FLOPS统计方法
存在这类方法,以下是基于TensorFlow内置工具的实现,分预测和训练阶段说明:
预测阶段(仅前向传播)
直接利用tf.profiler追踪模型前向传播的浮点运算量,代码示例:
import tensorflow as tf from tensorflow.keras import layers # 构建示例模型(替换为你的模型) model = tf.keras.Sequential([ layers.Dense(64, activation='relu', input_shape=(32,)), layers.Dense(10, activation='softmax') ]) # 生成匹配模型输入形状的张量(使用实际业务的batch size更准确) input_tensor = tf.random.uniform((1, 32)) # 启动profiler统计 with tf.profiler.Profile() as profiler: with tf.profiler.experimental.Trace('predict', step_num=1, _r=1): model(input_tensor) # 提取并打印FLOPS flops = profiler.profile_operations(options=tf.profiler.ProfileOptionBuilder.float_operation()) print(f"预测阶段单步FLOPS: {flops.total_float_ops}")
训练阶段(前向+反向传播)
训练阶段包含前向计算、损失计算和反向传播,需在完整训练步骤中追踪,代码示例:
# 定义训练依赖组件(替换为你的损失函数和优化器) loss_fn = tf.keras.losses.SparseCategoricalCrossentropy() optimizer = tf.keras.optimizers.Adam() # 模拟训练数据(替换为你的真实数据) x_sample = tf.random.uniform((1, 32)) y_sample = tf.random.uniform((1,), maxval=10, dtype=tf.int32) # 启动profiler统计单步训练FLOPS with tf.profiler.Profile() as profiler: with tf.profiler.experimental.Trace('train', step_num=1, _r=1): with tf.GradientTape() as tape: preds = model(x_sample) loss = loss_fn(y_sample, preds) grads = tape.gradient(loss, model.trainable_variables) optimizer.apply_gradients(zip(grads, model.trainable_variables)) # 提取并打印训练单步FLOPS train_flops = profiler.profile_operations(options=tf.profiler.ProfileOptionBuilder.float_operation()) print(f"训练单步FLOPS(含前向+反向传播): {train_flops.total_float_ops}")
注意事项
- 统计结果与输入的batch size直接相关,建议使用实际业务场景中的batch size以获得准确值
- CPU环境下无需额外配置,TensorFlow会自动适配CPU设备完成统计
- 上述代码统计的是单步计算量,若需统计整个训练周期的总FLOPS,需将单步值乘以总训练步数
内容的提问来源于stack exchange,提问作者Los
相关产品推荐
相关产品推荐

