TensorFlow两种操作数统计方法结果不一致,求技术解释
问题:TensorFlow两种FLOPs统计方法结果差异解析
统计TensorFlow模型操作数时,使用get_flops和get_flops_tfv2_1两种方法得到的结果分别为120和1440(后者接近理论计算值),以下是两种方法的统计逻辑差异解析:
模型代码
import tensorflow as tf from tensorflow.keras.models import Model from tensorflow.keras.layers import Input, Conv1D, GlobalAveragePooling1D def model_gen(win_size, n_feats): input1 = Input(shape=(win_size,n_feats), name='input') x = Conv1D(filters=120,kernel_size=1,dilation_rate=1, strides=1, padding='same')(input1) x13=GlobalAveragePooling1D()(x) model = Model(inputs=input1, outputs=x13) return model def get_flops(model_h5_path): tf.compat.v1.disable_eager_execution() session = tf.compat.v1.Session() graph = tf.compat.v1.get_default_graph() with graph.as_default(): with session.as_default(): model = tf.keras.models.load_model(model_h5_path) run_meta = tf.compat.v1.RunMetadata() opts = tf.compat.v1.profiler.ProfileOptionBuilder.float_operation() flops = tf.compat.v1.profiler.profile(graph=graph, run_meta=run_meta, cmd='op', options=opts) tf.compat.v1.reset_default_graph() return flops.total_float_ops def get_flops_tfv2_1(model_h5_path): from tensorflow.python.profiler.option_builder import ProfileOptionBuilder from tensorflow.python.profiler.model_analyzer import profile model = tf.keras.models.load_model(model_h5_path) forward_pass = tf.function(model.call, input_signature=[tf.TensorSpec(shape=(1,) + model.input_shape[1:])]) graph_info = profile(forward_pass.get_concrete_function().graph, options=ProfileOptionBuilder.float_operation()) # The //2 is necessary since `profile` counts multiply and accumulate # as two flops, here we report the total number of multiply accumulate ops flops = graph_info.total_float_ops // 2 return flops if __name__ == '__main__': dur = 6 model_file = 'tmp.h5' n_feats=1 model=model_gen(dur, n_feats) model.save(model_file) ops = get_flops(model_file) # ops = get_flops_tfv2_1(model_file) print(f'#Operations = {ops}')
两种方法的统计逻辑差异
1. get_flops(TF1.x兼容模式)
该方法基于TensorFlow 1.x的静态图机制,通过禁用eager execution切换到旧版环境。但TF1.x的profiler在加载Keras模型时存在局限性:
- 仅统计了模型中GlobalAveragePooling1D层的部分操作(返回的120正好是池化层的输出特征数),完全遗漏了Conv1D层的核心浮点操作。
- 原因是旧版profiler无法完整追踪Keras模型转换后的静态计算图,导致卷积层的计算节点未被统计。
2. get_flops_tfv2_1(TF2.x原生方法)
该方法利用TF2.x的tf.function将模型转换为完整的静态计算图,新版profiler能准确遍历所有计算节点:
- 完整统计Conv1D层的浮点操作:每个输出元素包含1次乘法(输入特征×卷积核)+1次加法(加偏置),针对
6个时间步×120个滤波器的输出,总操作数为6×120×2=1440,与理论计算值完全匹配。 - 代码中的
//2注释是通用场景下的处理(将乘加操作合并计数),但在当前模型的统计中,profiler直接返回了乘加操作的总和,因此结果正好对应理论值。
内容的提问来源于stack exchange,提问作者Vinay
相关产品推荐
相关产品推荐

