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

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.05.07 21:53:02