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

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.24 09:57:18