TensorFlow1.15中Keras模型FLOPS计算因形状不全结果不准的解决方法
解决TensorFlow 1.15中Keras模型FLOPS统计不准的问题
你遇到的核心问题是动态batch_size(即shape中的None)导致TensorFlow Profiler无法确定部分操作的张量形状,进而无法统计这些操作的FLOPS,最终结果偏差较大。结合你的TF1.15环境,这里提供两种可行的解决方案:
方法一:通过固定输入触发前向传播,确定所有张量形状
在统计FLOPS前,让模型处理一个固定形状的输入,迫使TensorFlow将图中所有操作的张量形状具体化,Profiler就能准确统计全量FLOPS了。修改后的代码如下:
import tensorflow as tf import numpy as np # 1. 创建ResNet50模型,明确输入尺寸(不含batch_size) model = tf.keras.applications.ResNet50( include_top=True, weights="imagenet", input_tensor=None, input_shape=(224, 224, 3), pooling=None, classes=1000 ) # 2. 生成固定batch_size的输入(这里取batch_size=1,匹配ResNet50默认输入尺寸) fixed_input = tf.random.normal((1, 224, 224, 3)) # 3. 运行一次前向传播,让TensorFlow确定所有张量的具体形状 with tf.compat.v1.Session() as sess: sess.run(tf.compat.v1.global_variables_initializer()) _ = sess.run(model(fixed_input)) # 4. 统计可训练参数数量 nparams = np.sum([np.prod(v.get_shape().as_list()) for v in tf.compat.v1.trainable_variables()]) print(f"可训练参数数量: {nparams}") # 5. 统计FLOPS options = tf.profiler.ProfileOptionBuilder.float_operation() options['output'] = 'none' flops = tf.profiler.profile(tf.get_default_graph(), options=options).total_float_ops # 除以2是因为Profiler会把乘加操作拆成两次浮点运算统计,而通常我们将乘加合并为一次计算FLOPS flops = flops // 2 print(f"模型FLOPS: {flops}")
方法二:创建模型时直接指定固定形状的输入张量
另一种思路是在构造模型时传入固定形状的占位符作为input_tensor,让模型从一开始就使用确定的张量形状,无需额外运行前向传播:
import tensorflow as tf import numpy as np # 1. 创建固定形状的输入占位符(明确batch_size=1) input_tensor = tf.compat.v1.placeholder(tf.float32, shape=(1, 224, 224, 3)) # 2. 基于固定输入张量创建ResNet50模型 model = tf.keras.applications.ResNet50( include_top=True, weights="imagenet", input_tensor=input_tensor, pooling=None, classes=1000 ) # 3. 统计参数和FLOPS nparams = np.sum([np.prod(v.get_shape().as_list()) for v in tf.compat.v1.trainable_variables()]) print(f"可训练参数数量: {nparams}") options = tf.profiler.ProfileOptionBuilder.float_operation() options['output'] = 'none' flops = tf.profiler.profile(tf.get_default_graph(), options=options).total_float_ops flops = flops // 2 print(f"模型FLOPS: {flops}")
关键说明
- 两种方法的核心都是消除张量形状中的动态维度(
None),让Profiler能够精准计算每一个操作的FLOPS,使用后你之前遇到的111 ops no flops stats due to incomplete shapes提示会消失。 - 统计出的FLOPS会和ResNet50的理论值(约4.1×10^9次浮点运算)接近,结果具备参考性。
- 注意在TF1.15环境下,需使用
tf.compat.v1下的Session及相关API,因为默认采用图模式运行。
内容的提问来源于stack exchange,提问作者Eypros
相关产品推荐
相关产品推荐

